"""Original SARIMAX study audit. Python 3.12; no downloads, no file writes.

python audit.py /path/to/inflation-forecast > audit-results.json
Optional --auto reruns the original full-sample stepwise search (not validation).
The input snapshot is NOT independently certified official data. This script
does not redistribute it, impute missing observations, or suppress warnings.
"""

import argparse
import hashlib
import importlib.metadata
import json
from pathlib import Path

import numpy as np
import pandas as pd
from sklearn.model_selection import TimeSeriesSplit
from statsmodels.stats.diagnostic import acorr_ljungbox
from statsmodels.tsa.statespace.sarimax import SARIMAX
from statsmodels.tsa.stattools import adfuller

TARGET = "Enflasyon_Aylik_Yuzde"
FEATURES = ["Kur_Degisim", "Petrol_Degisim", "Faiz_Degisim"]


def prepare(path):
    raw = pd.read_csv(path, parse_dates=["Tarih"]).set_index("Tarih")
    if not raw.index.is_unique or not raw.index.is_monotonic_increasing:
        raise ValueError("Dates must be unique and sorted; no silent repair.")
    expected = pd.date_range(raw.index.min(), raw.index.max(), freq="MS")
    if not raw.index.equals(expected):
        raise ValueError("Missing months or non-month-start timestamps.")
    if not np.isfinite(raw.to_numpy(dtype=float)).all():
        raise ValueError("Missing/nonfinite raw values: verify sources first.")
    raw = raw.asfreq("MS")
    data = raw.assign(
        Kur_Degisim=raw.USD_TRY.pct_change(fill_method=None) * 100,
        Petrol_Degisim=raw.Brent_Petrol_USD.pct_change(fill_method=None) * 100,
        Faiz_Degisim=raw.Faiz_Orani.diff(),  # percentage POINTS, not pct_change
    ).iloc[1:]
    return raw, data


def fit(y, x, maxiter=50):
    # Fixed original specification, including its no-trend refit.
    return SARIMAX(
        y, exog=x, order=(0, 1, 2), seasonal_order=(0, 0, 0, 12),
        enforce_stationarity=False, enforce_invertibility=False,
    ).fit(disp=False, maxiter=maxiter)


def future_scenario(x_train, dates):
    # Uses training history ONLY. A scenario, NOT a forecast of the predictors.
    values = [x_train.Kur_Degisim.tail(6).mean(),
              x_train.Petrol_Degisim.tail(6).mean(), 0.0]
    return pd.DataFrame(np.tile(values, (len(dates), 1)),
                        index=dates, columns=FEATURES)


def score(actual, predicted):
    error = np.asarray(actual) - np.asarray(predicted)
    return {"mae_pp": float(np.abs(error).mean()),
            "rmse_pp": float(np.sqrt(np.square(error).mean()))}


def evaluate(data, horizon=None):
    y, x = data[TARGET], data[FEATURES]
    # Legacy split for reproduction; same origins with 17-month tests separately.
    splits = TimeSeriesSplit(n_splits=3, test_size=horizon)
    rows = []
    for fold, (train, test) in enumerate(splits.split(data), start=1):
        y_train, y_test = y.iloc[train], y.iloc[test]
        x_train, x_test = x.iloc[train], x.iloc[test]
        fitted = fit(y_train, x_train)
        forecasts = {
            "oracle_future_x": fitted.get_forecast(len(test), exog=x_test).predicted_mean,
            "training_only_scenario": fitted.get_forecast(
                len(test), exog=future_scenario(x_train, y_test.index)
            ).predicted_mean,
            "last_value": np.repeat(y_train.iloc[-1], len(test)),
            "training_mean": np.repeat(y_train.mean(), len(test)),
            "seasonal_naive": np.resize(y_train.iloc[-12:].to_numpy(), len(test)),
        }
        rows.append({
            "fold": fold, "train_n": len(train), "test_n": len(test),
            "train_end": str(y_train.index[-1].date()),
            "test_start": str(y_test.index[0].date()),
            "test_end": str(y_test.index[-1].date()),
            "converged": bool(fitted.mle_retvals.get("converged")),
            "metrics": {name: score(y_test, pred) for name, pred in forecasts.items()},
        })
    means = {
        name: {metric: float(np.mean([r["metrics"][name][metric] for r in rows]))
               for metric in ("mae_pp", "rmse_pp")}
        for name in rows[0]["metrics"]
    }
    return {"folds": rows, "mean_fold_metrics": means,
            "caveat": "Fixed order selected on full sample in original study; "
                      "not nested validation. No vintage/release-lag data available."}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("repository", type=Path)
    parser.add_argument("--auto", action="store_true")
    args = parser.parse_args()
    path = args.repository / "makro_veri_seti_kapsamli.csv"
    raw, data = prepare(path)
    y, x = data[TARGET], data[FEATURES]
    model = fit(y, x)
    forecast_dates = pd.date_range(data.index[-1] + pd.offsets.MonthBegin(1), periods=17, freq="MS")
    forecast = model.get_forecast(17, exog=future_scenario(x, forecast_dates))
    intervals = forecast.conf_int(alpha=0.05)
    cpi_error = (raw.TUFE_Endeks.pct_change(fill_method=None) * 100 - raw[TARGET]).dropna()
    residuals = model.resid.iloc[model.loglikelihood_burn:]
    report = {
        "source_commit": "79c01d5173322e6e25f7c95e88d83aeb810fc021",
        "csv_sha256": hashlib.sha256(path.read_bytes()).hexdigest(),
        "versions": {p: importlib.metadata.version(p) for p in (
            "numpy", "pandas", "scipy", "statsmodels", "scikit-learn", "pmdarima")},
        "data": {"raw_n": len(raw), "model_n": len(data),
                 "start": str(raw.index[0].date()), "end": str(raw.index[-1].date()),
                 "cpi_monthly_max_abs_difference_pp": float(cpi_error.abs().max()),
                 "cpi_mismatches_over_001_pp": {str(k.date()): float(v) for k, v in cpi_error[cpi_error.abs() > .01].items()},
                 "provenance": "May 2025–July 2026 manually appended; original release URLs/vintages absent."},
        "target_correlations": data[[TARGET] + FEATURES].corr()[TARGET].to_dict(),
        "adf_pvalues": {col: float(adfuller(data[col])[1]) for col in [TARGET] + FEATURES},
        "full_model": {"order": [0, 1, 2], "seasonal_order": [0, 0, 0, 12],
                       "converged": bool(model.mle_retvals.get("converged")),
                       "aic": float(model.aic), "burn_in": model.loglikelihood_burn,
                       "params": model.params.to_dict(),
                       "lb_original_lag12_p": float(acorr_ljungbox(model.resid, lags=[12], return_df=True).lb_pvalue.iloc[0]),
                       "lb_burn_adjusted_df2_lag12_p": float(acorr_ljungbox(residuals, lags=[12], model_df=2, return_df=True).lb_pvalue.iloc[0])},
        "legacy_cv": evaluate(data),
        "horizon17_cv": evaluate(data, horizon=17),
        "conditional_forecast": {
            "future_predictors_each_month": future_scenario(x, forecast_dates).iloc[0].to_dict(),
            "points": [{"date": str(date.date()), "mean_pct": float(forecast.predicted_mean.iloc[i]),
                        "lower95_pct": float(intervals.iloc[i, 0]), "upper95_pct": float(intervals.iloc[i, 1])}
                       for i, date in enumerate(forecast_dates)],
            "caveat": "Conditional on fixed exogenous paths/model. Not all-shock coverage; not annual inflation."
        },
    }
    if args.auto:
        import pmdarima as pm
        auto = pm.auto_arima(y, X=x, seasonal=True, m=12, stepwise=True,
                             suppress_warnings=False, error_action="warn")
        report["original_auto_search_rerun"] = {
            "order": auto.order, "seasonal_order": auto.seasonal_order,
            "with_intercept": auto.with_intercept, "aic": auto.aic(),
            "converged": bool(auto.arima_res_.mle_retvals.get("converged")),
            "caveat": "Full-sample search, only for reproduction; original environment unpinned."
        }
    print(json.dumps(report, ensure_ascii=False, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
