"""
Train SARIMAX market price forecasting models — one per (district, crop) pair.

Usage:
    python -m ml_pipeline.training.train_sarimax_market \
        --input data/processed/market_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 = 24  # market series can be shorter than climate's 36-month
                            # requirement since price data tends to be denser
                            # and seasonal patterns shift faster than climate


def train_pair_model(key: str, series: pd.Series) -> tuple[object | None, dict]:
    if len(series) < MIN_MONTHS_REQUIRED:
        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=1, max_Q=1,
            )

        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 '{key}': {exc}")
        return None, {"status": "failed", "reason": str(exc)}


def main() -> None:
    parser = argparse.ArgumentParser(description="Train SARIMAX market price models per district-crop pair")
    parser.add_argument("--input", required=True)
    parser.add_argument("--output-dir", required=True)
    parser.add_argument(
        "--min-pair-frequency", type=int, default=24,
        help="Skip district-crop pairs with fewer than this many monthly observations",
    )
    args = parser.parse_args()

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

    df["pair_key"] = df["district"] + "_" + df["crop"]
    pair_counts = df["pair_key"].value_counts()
    eligible_pairs = sorted(pair_counts[pair_counts >= args.min_pair_frequency].index)

    logger.info(
        f"{len(eligible_pairs)} of {df['pair_key'].nunique()} district-crop pairs "
        f"have enough history (>= {args.min_pair_frequency} months) to train on"
    )

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

    for key in eligible_pairs:
        d = df[df["pair_key"] == key].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["price_per_quintal"].asfreq("MS")
        series = series.interpolate()  # SARIMAX requires no internal gaps in the index

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

    if not models:
        raise RuntimeError(
            "Zero district-crop market models were successfully trained. "
            "Check that market_clean.csv has enough monthly history per pair."
        )

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

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

    save_training_metadata(
        output_dir,
        model_name="sarimax_market",
        metrics={"avg_aic": avg_aic, "n_pairs_trained": len(models), "per_pair_report": training_report},
        feature_cols=["price_per_quintal (univariate time series per district-crop pair)"],
        n_train=sum(r.get("n_months", 0) for r in training_report.values()),
        n_test=0,
    )

    logger.info(f"Trained {len(models)}/{len(eligible_pairs)} district-crop models successfully")
    logger.info(f"Saved to {output_path}")
    logger.info(
        "Pairs without a trained model will return degraded market results "
        "in the live API (see app/services/model_runners.run_market_model)."
    )


if __name__ == "__main__":
    main()
