"""
Train SARIMAX climate forecasting models — one per district.

Usage:
    python -m ml_pipeline.training.train_sarimax_climate \
        --input data/processed/climate_clean.csv \
        --output-dir data/models/v1
"""

import argparse
import warnings
from pathlib import Path

import joblib
import pandas as pd
import pmdarima as pm
from statsmodels.tsa.statespace.sarimax import SARIMAX

from ml_pipeline.training.training_utils import get_logger, save_training_metadata

logger = get_logger(__name__)

MIN_MONTHS_REQUIRED = 36  # need at least 3 years of monthly data per district


def train_district_model(district: str, series: pd.Series) -> tuple[object | None, dict]:
    if len(series) < MIN_MONTHS_REQUIRED:
        logger.warning(
            f"District '{district}' has only {len(series)} months of data "
            f"(need {MIN_MONTHS_REQUIRED}+) — skipping"
        )
        return None, {"status": "skipped", "reason": "insufficient_data", "n_months": len(series)}

    try:
        with warnings.catch_warnings():
            warnings.simplefilter("ignore")
            auto = pm.auto_arima(
                series,
                seasonal=True,
                m=12,
                stepwise=True,
                suppress_warnings=True,
                error_action="ignore",
                max_p=3, max_q=3, max_P=2, max_Q=2,
            )

        p, d, q = auto.order
        P, D, Q, s = auto.seasonal_order

        model = SARIMAX(
            series,
            order=(p, d, q),
            seasonal_order=(P, D, Q, s),
            enforce_stationarity=False,
            enforce_invertibility=False,
        )
        result = model.fit(disp=False)

        return result, {
            "status": "trained",
            "order": [p, d, q],
            "seasonal_order": [P, D, Q, s],
            "aic": float(result.aic),
            "n_months": len(series),
        }
    except Exception as exc:
        logger.error(f"Training failed for district '{district}': {exc}")
        return None, {"status": "failed", "reason": str(exc)}


def main() -> None:
    parser = argparse.ArgumentParser(description="Train SARIMAX climate models per district")
    parser.add_argument("--input", required=True)
    parser.add_argument("--output-dir", required=True)
    args = parser.parse_args()

    df = pd.read_csv(args.input)
    required_cols = {"district", "year", "month", "rainfall_mm"}
    missing = required_cols - set(df.columns)
    if missing:
        raise ValueError(f"Input file missing required columns: {missing}")

    districts = sorted(df["district"].unique())
    logger.info(f"Training SARIMAX climate models for {len(districts)} districts")

    models: dict[str, object] = {}
    training_report: dict[str, dict] = {}

    for district in districts:
        d = df[df["district"] == district].sort_values(["year", "month"]).copy()
        d["date"] = pd.to_datetime(d["year"].astype(str) + "-" + d["month"].astype(str) + "-01")
        d = d.set_index("date")
        series = d["rainfall_mm"].asfreq("MS")  # month-start frequency, required by SARIMAX

        model, report = train_district_model(district, series)
        training_report[district] = report
        if model is not None:
            models[district] = model

    if not models:
        raise RuntimeError(
            "Zero district models were successfully trained. Check that "
            "climate_clean.csv has enough history per district (36+ months)."
        )

    output_dir = Path(args.output_dir)
    output_dir.mkdir(parents=True, exist_ok=True)
    output_path = output_dir / "sarimax_climate.pkl"
    joblib.dump(models, output_path)

    avg_aic = sum(
        r["aic"] for r in training_report.values() if r.get("status") == "trained"
    ) / max(1, sum(1 for r in training_report.values() if r.get("status") == "trained"))

    save_training_metadata(
        output_dir,
        model_name="sarimax_climate",
        metrics={"avg_aic": avg_aic, "per_district_report": training_report},
        feature_cols=["rainfall_mm (univariate time series per district)"],
        n_train=sum(r.get("n_months", 0) for r in training_report.values()),
        n_test=0,  # SARIMAX validated via AIC and walk-forward, not a held-out split here
    )

    logger.info(f"Trained {len(models)}/{len(districts)} district models successfully")
    logger.info(f"Saved to {output_path}")
    skipped = [d for d, r in training_report.items() if r["status"] != "trained"]
    if skipped:
        logger.warning(
            f"{len(skipped)} districts have NO trained climate model: {skipped}. "
            f"The live API will return degraded climate results for these "
            f"districts (see app/services/model_runners.run_climate_model)."
        )


if __name__ == "__main__":
    main()
