"""Reproducible leakage-safe baseline and temporal-memory experiment."""

from __future__ import annotations

import hashlib
import json
import os
import platform
import random
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import joblib
import sklearn
import torch
import xgboost
import yaml
from sklearn.impute import SimpleImputer
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score
from sklearn.preprocessing import StandardScaler
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from xgboost import XGBRegressor

FEATURES = [
    "manure_fed_kg", "water_added_kg", "air_temperature_c",
    "target_lag_1", "target_roll_mean_3", "target_roll_std_3",
    "day_sin", "day_cos", "reactor_R1", "reactor_R2", "reactor_R3", "reactor_R4",
]
TARGET = "total_biogas_interval_ml"


def _sha(path: Path) -> str:
    h = hashlib.sha256()
    with path.open("rb") as f:
        for b in iter(lambda: f.read(1024 * 1024), b""):
            h.update(b)
    return h.hexdigest()


def _git_commit(repo: Path) -> str:
    return subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=repo, text=True).strip()


def set_seed(seed: int) -> None:
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.use_deterministic_algorithms(True)


def load_features(csv_path: Path) -> pd.DataFrame:
    df = pd.read_csv(csv_path, parse_dates=["timestamp"])
    # A source-labelled "daily" increment is only comparable for one-day
    # intervals. Duplicate dates, discontinuities, and missing targets are
    # excluded, never imputed.
    df = df[
        df[TARGET].notna()
        & ~df["duplicate_timestamp"].astype(bool)
        & df["interval_days"].eq(1)
    ].copy()
    df = df.sort_values(["reactor_id", "timestamp"], kind="stable")
    grouped = df.groupby("reactor_id", sort=False)[TARGET]
    df["target_lag_1"] = grouped.shift(1)
    df["target_roll_mean_3"] = grouped.transform(lambda s: s.shift(1).rolling(3).mean())
    df["target_roll_std_3"] = grouped.transform(lambda s: s.shift(1).rolling(3).std())
    day = df["timestamp"].dt.dayofyear
    df["day_sin"] = np.sin(2 * np.pi * day / 366)
    df["day_cos"] = np.cos(2 * np.pi * day / 366)
    for reactor in ("R1", "R2", "R3", "R4"):
        df[f"reactor_{reactor}"] = (df["reactor_id"] == reactor).astype(float)
    return df.dropna(subset=["target_lag_1"]).reset_index(drop=True)


def make_split_manifest(df: pd.DataFrame) -> dict:
    dates = np.array(sorted(df["timestamp"].unique()))
    n = len(dates)
    test_i = int(np.floor(n * 0.85))
    test_start = pd.Timestamp(dates[test_i])
    pretest = dates[:test_i]
    boundaries = [0.55, 0.65, 0.75, 0.85]
    cut = [pd.Timestamp(dates[min(int(np.floor(n * x)), n - 1)]) for x in boundaries]
    folds = []
    for i in range(3):
        folds.append({
            "fold": i + 1,
            "train_end_exclusive": str(cut[i].date()),
            "validation_start": str(cut[i].date()),
            "validation_end_exclusive": str(cut[i + 1].date()),
        })
    return {
        "version": "ds03-chronological-v1",
        "group_key": "reactor_id",
        "date_count": n,
        "pretest_end_exclusive": str(test_start.date()),
        "final_test_start": str(test_start.date()),
        "final_test_end": str(pd.Timestamp(dates[-1]).date()),
        "final_test_untouched_during_selection": True,
        "rolling_folds": folds,
    }


def metrics(y: np.ndarray, pred: np.ndarray) -> dict:
    denom = np.abs(y) + np.abs(pred)
    valid = denom > 1e-9
    return {
        "mae": float(mean_absolute_error(y, pred)),
        "rmse": float(mean_squared_error(y, pred) ** 0.5),
        "r2": float(r2_score(y, pred)),
        "smape": float(np.mean(200 * np.abs(pred[valid] - y[valid]) / denom[valid])) if valid.any() else None,
        "n": int(len(y)),
    }


def sequences(df: pd.DataFrame, window: int) -> tuple[np.ndarray, np.ndarray, pd.DataFrame]:
    xs, ys, meta = [], [], []
    for _, part in df.groupby("reactor_id", sort=True):
        part = part.sort_values("timestamp").reset_index(drop=True)
        values = part[FEATURES].to_numpy(float)
        target = part[TARGET].to_numpy(float)
        for i in range(window - 1, len(part)):
            xs.append(values[i - window + 1:i + 1])
            ys.append(target[i])
            meta.append(part.loc[i, ["timestamp", "reactor_id", "air_temperature_c"]].to_dict())
    return np.asarray(xs), np.asarray(ys), pd.DataFrame(meta)


class BaselineLSTM(nn.Module):
    def __init__(self, n_features: int, hidden: int, layers: int):
        super().__init__()
        self.lstm = nn.LSTM(n_features, hidden, num_layers=layers, batch_first=True)
        self.head = nn.Linear(hidden, 1)

    def forward(self, x):
        values, _ = self.lstm(x)
        return self.head(values[:, -1]).squeeze(-1)


def _fit_transform(train_x, val_x):
    shape_t, shape_v = train_x.shape, val_x.shape
    imp = SimpleImputer(strategy="median").fit(train_x.reshape(-1, shape_t[-1]))
    tr = imp.transform(train_x.reshape(-1, shape_t[-1]))
    va = imp.transform(val_x.reshape(-1, shape_v[-1]))
    scale = StandardScaler().fit(tr)
    return scale.transform(tr).reshape(shape_t), scale.transform(va).reshape(shape_v)


def fit_lstm(train_x, train_y, val_x, val_y, cfg, seed, epoch_rows):
    set_seed(seed)
    train_x, val_x = _fit_transform(train_x, val_x)
    y_scaler = StandardScaler().fit(train_y.reshape(-1, 1))
    ty = y_scaler.transform(train_y.reshape(-1, 1)).ravel()
    vy = y_scaler.transform(val_y.reshape(-1, 1)).ravel()
    model = BaselineLSTM(train_x.shape[-1], cfg["hidden_size"], cfg["num_layers"])
    optimizer = torch.optim.Adam(model.parameters(), lr=cfg["learning_rate"])
    loss_fn = nn.MSELoss()
    loader = DataLoader(
        TensorDataset(torch.tensor(train_x, dtype=torch.float32), torch.tensor(ty, dtype=torch.float32)),
        batch_size=cfg["batch_size"], shuffle=False,
    )
    vx = torch.tensor(val_x, dtype=torch.float32)
    vy_t = torch.tensor(vy, dtype=torch.float32)
    best, state, patience = float("inf"), None, 0
    for epoch in range(1, cfg["epochs"] + 1):
        model.train()
        losses = []
        for bx, by in loader:
            optimizer.zero_grad()
            loss = loss_fn(model(bx), by)
            loss.backward()
            optimizer.step()
            losses.append(float(loss.detach()))
        model.eval()
        with torch.no_grad():
            val_loss = float(loss_fn(model(vx), vy_t))
        epoch_rows.append({"epoch": epoch, "train_loss": np.mean(losses), "validation_loss": val_loss})
        if val_loss < best - 1e-7:
            best, state, patience = val_loss, {k: v.detach().clone() for k, v in model.state_dict().items()}, 0
        else:
            patience += 1
            if patience >= cfg["patience"]:
                break
    model.load_state_dict(state)
    model.eval()
    with torch.no_grad():
        pred = model(vx).numpy()
    return y_scaler.inverse_transform(pred.reshape(-1, 1)).ravel(), model


def drift_metrics(train_temp, val_temp, y, pred, low_q, high_q):
    low, high = np.nanquantile(train_temp, [low_q, high_q])
    drift = (val_temp < low) | (val_temp > high)
    result = {"training_temperature_low_c": float(low), "training_temperature_high_c": float(high),
              "natural_drift_n": int(drift.sum()), "stable_n": int((~drift).sum())}
    if drift.sum() >= 2:
        result["drift"] = metrics(y[drift], pred[drift])
    if (~drift).sum() >= 2:
        result["stable"] = metrics(y[~drift], pred[~drift])
    if "drift" in result and "stable" in result:
        result["rmse_degradation_percent"] = float(
            100 * (result["drift"]["rmse"] - result["stable"]["rmse"]) / max(result["stable"]["rmse"], 1e-9)
        )
    return result


def run(repo: Path, csv_path: Path, config_path: Path, manifest_path: Path, run_root: Path, run_id: str) -> Path:
    cfg = yaml.safe_load(config_path.read_text())
    df = load_features(csv_path)
    split = make_split_manifest(df)
    run_dir = run_root / run_id
    if run_dir.exists():
        raise FileExistsError(f"Run already exists: {run_dir}")
    for name in ("logs", "predictions", "checkpoints", "plots"):
        (run_dir / name).mkdir(parents=True, exist_ok=True)
    (run_dir / "config.yaml").write_text(config_path.read_text())
    (run_dir / "split_manifest.json").write_text(json.dumps(split, indent=2) + "\n")
    processing = json.loads(manifest_path.read_text())
    environment = {
        "python": sys.version, "platform": platform.platform(),
        "numpy": np.__version__, "pandas": pd.__version__,
        "scikit_learn": sklearn.__version__, "xgboost": xgboost.__version__,
        "torch": torch.__version__,
    }
    (run_dir / "environment.json").write_text(json.dumps(environment, indent=2) + "\n")

    all_metrics, all_predictions, all_epochs = [], [], []
    fold_defs = split["rolling_folds"]
    for window in cfg["sequence_windows"]:
        x, y, meta = sequences(df, window)
        for fold in fold_defs:
            train_end = pd.Timestamp(fold["train_end_exclusive"])
            val_start = pd.Timestamp(fold["validation_start"])
            val_end = pd.Timestamp(fold["validation_end_exclusive"])
            tr = meta["timestamp"] < train_end
            va = (meta["timestamp"] >= val_start) & (meta["timestamp"] < val_end)
            if tr.sum() < 20 or va.sum() < 5:
                continue
            # Persistence has no fitted parameters.
            persistence = x[va, -1, FEATURES.index("target_lag_1")]
            for model_name, pred in [("persistence", persistence)]:
                row = {"model": model_name, "window": window, "fold": fold["fold"], "seed": None, **metrics(y[va], pred)}
                all_metrics.append(row)
                for m, actual, estimate in zip(meta.loc[va].to_dict("records"), y[va], pred):
                    all_predictions.append({**m, "model": model_name, "window": window, "fold": fold["fold"],
                                            "seed": None, "actual": actual, "prediction": estimate})
            for seed in cfg["seeds"]:
                # Fold-local median imputation; XGBoost receives only past-only
                # lags/rolling features plus contemporaneously observable inputs.
                train_tab, val_tab = x[tr, -1], x[va, -1]
                imp = SimpleImputer(strategy="median").fit(train_tab)
                model = XGBRegressor(random_state=seed, n_jobs=1, objective="reg:squarederror", **cfg["xgboost"])
                model.fit(imp.transform(train_tab), y[tr])
                pred = model.predict(imp.transform(val_tab))
                all_metrics.append({"model": "xgboost", "window": window, "fold": fold["fold"], "seed": seed, **metrics(y[va], pred)})
                for m, actual, estimate in zip(meta.loc[va].to_dict("records"), y[va], pred):
                    all_predictions.append({**m, "model": "xgboost", "window": window, "fold": fold["fold"],
                                            "seed": seed, "actual": actual, "prediction": float(estimate)})
                epoch_rows = []
                pred, model_lstm = fit_lstm(x[tr], y[tr], x[va], y[va], cfg["lstm"], seed, epoch_rows)
                for e in epoch_rows:
                    all_epochs.append({**e, "model": "lstm", "window": window, "fold": fold["fold"], "seed": seed})
                all_metrics.append({"model": "lstm", "window": window, "fold": fold["fold"], "seed": seed, **metrics(y[va], pred)})
                for m, actual, estimate in zip(meta.loc[va].to_dict("records"), y[va], pred):
                    all_predictions.append({**m, "model": "lstm", "window": window, "fold": fold["fold"],
                                            "seed": seed, "actual": actual, "prediction": float(estimate)})

    metric_df = pd.DataFrame(all_metrics)
    pred_df = pd.DataFrame(all_predictions)
    epoch_df = pd.DataFrame(all_epochs)
    metric_df.to_csv(run_dir / "fold_metrics.csv", index=False)
    pred_df.to_csv(run_dir / "predictions" / "rolling_predictions.csv", index=False)
    epoch_df.to_csv(run_dir / "logs" / "epoch_history.csv", index=False)

    summary = (
        metric_df.groupby(["model", "window"])[["mae", "rmse", "r2", "smape"]]
        .agg(["mean", "std"]).reset_index()
    )
    summary.columns = ["_".join(str(x) for x in c if x != "") for c in summary.columns]
    summary.to_csv(run_dir / "rolling_summary.csv", index=False)
    # Robust selection: collapse seeds, rank within each fold, and select the
    # model with most fold wins. This prevents one extreme fold from making a
    # high-variance model look universally superior.
    fold_model = metric_df.groupby(["model", "window", "fold"], as_index=False)["rmse"].mean()
    window_robust = (
        fold_model.groupby(["model", "window"], as_index=False)["rmse"]
        .agg(rmse_median="median", rmse_mean="mean", rmse_std="std")
    )
    best_windows = window_robust.sort_values(
        ["model", "rmse_median", "rmse_std"]
    ).groupby("model", as_index=False).first()
    candidates = fold_model.merge(best_windows[["model", "window"]], on=["model", "window"])
    winners = candidates.loc[candidates.groupby("fold")["rmse"].idxmin()]
    wins = winners["model"].value_counts().to_dict()
    best_windows["fold_wins"] = best_windows["model"].map(wins).fillna(0)
    chosen = best_windows.sort_values(
        ["fold_wins", "rmse_median", "rmse_std"],
        ascending=[False, True, True],
    ).iloc[0]
    selected_model, selected_window = str(chosen["model"]), int(chosen["window"])

    selected_preds = pred_df[(pred_df.model == selected_model) & (pred_df.window == selected_window)].copy()
    drift_rows = []
    for (fold, seed), part in selected_preds.groupby(["fold", "seed"], dropna=False):
        train_end = pd.Timestamp(fold_defs[int(fold) - 1]["train_end_exclusive"])
        train_temp = df.loc[df.timestamp < train_end, "air_temperature_c"].dropna().to_numpy()
        val_temp = part["air_temperature_c"].to_numpy(float)
        drift_rows.append({"fold": int(fold), "seed": None if pd.isna(seed) else int(seed),
                           **drift_metrics(train_temp, val_temp, part.actual.to_numpy(), part.prediction.to_numpy(),
                                           cfg["temperature_drift"]["low_quantile"], cfg["temperature_drift"]["high_quantile"])})
    (run_dir / "temperature_drift.json").write_text(json.dumps(drift_rows, indent=2) + "\n")

    # Open the locked final period exactly once, after model/window selection.
    x, y, meta = sequences(df, selected_window)
    test_start = pd.Timestamp(split["final_test_start"])
    tr = meta["timestamp"] < test_start
    te = meta["timestamp"] >= test_start
    final_metrics, final_predictions = [], []
    final_models = ["persistence", "xgboost", "lstm"]
    for model_name in final_models:
        seeds = [None] if model_name == "persistence" else cfg["seeds"]
        for seed in seeds:
            if model_name == "persistence":
                pred = x[te, -1, FEATURES.index("target_lag_1")]
            elif model_name == "xgboost":
                imp = SimpleImputer(strategy="median").fit(x[tr, -1])
                fitted = XGBRegressor(
                    random_state=seed, n_jobs=1, objective="reg:squarederror", **cfg["xgboost"]
                )
                fitted.fit(imp.transform(x[tr, -1]), y[tr])
                pred = fitted.predict(imp.transform(x[te, -1]))
                joblib.dump(
                    {"model": fitted, "imputer": imp, "features": FEATURES},
                    run_dir / "checkpoints" / f"xgboost_seed_{seed}.joblib",
                )
            else:
                final_epochs = []
                pred, fitted = fit_lstm(x[tr], y[tr], x[te], y[te], cfg["lstm"], seed, final_epochs)
                for e in final_epochs:
                    all_epochs.append({**e, "model": "lstm_final", "window": selected_window,
                                       "fold": "final_test", "seed": seed})
                torch.save(fitted.state_dict(), run_dir / "checkpoints" / f"lstm_seed_{seed}.pt")
            final_metrics.append({"model": model_name, "window": selected_window, "seed": seed, **metrics(y[te], pred)})
            for m, actual, estimate in zip(meta.loc[te].to_dict("records"), y[te], pred):
                final_predictions.append({**m, "model": model_name, "window": selected_window,
                                          "seed": seed, "actual": actual, "prediction": float(estimate)})
    pd.DataFrame(final_metrics).to_csv(run_dir / "final_test_metrics.csv", index=False)
    pd.DataFrame(final_predictions).to_csv(run_dir / "predictions" / "final_test_predictions.csv", index=False)
    # Rewrite epoch history to include locked final fits.
    pd.DataFrame(all_epochs).to_csv(run_dir / "logs" / "epoch_history.csv", index=False)

    plt.figure(figsize=(8, 4.5))
    for model, part in summary.groupby("model"):
        plt.errorbar(part["window"], part["rmse_mean"], yerr=part["rmse_std"], marker="o", label=model)
    plt.xlabel("Sequence window (observations)"); plt.ylabel("Rolling-origin RMSE (mL)")
    plt.title("GFIS Gate 3 temporal-memory study"); plt.legend(); plt.tight_layout()
    plt.savefig(run_dir / "plots" / "window_rmse.png", dpi=180); plt.close()

    decision = {
        "selection_basis": "most rolling-fold wins, then median fold RMSE and dispersion; selection completed before final test was opened",
        "selected_model": selected_model,
        "selected_window": selected_window,
        "fold_wins": wins,
        "rolling_rmse_median": float(chosen["rmse_median"]),
        "rolling_rmse_mean": float(chosen["rmse_mean"]),
        "rolling_rmse_std": float(chosen["rmse_std"]) if not pd.isna(chosen["rmse_std"]) else None,
        "claim_temporal_learning_improved": selected_model == "lstm",
        "lstm_status": "challenger only; not accepted as champion" if selected_model != "lstm" else "selected",
        "physics_violation_rate": None,
        "physics_violation_reason": cfg["physics_violation"]["reason"],
        "final_test_status": "opened once after architecture selection; see final_test_metrics.csv",
    }
    (run_dir / "scientific_decision.json").write_text(json.dumps(decision, indent=2) + "\n")

    artifacts = {}
    for path in sorted(p for p in run_dir.rglob("*") if p.is_file()):
        artifacts[str(path.relative_to(run_dir))] = _sha(path)
    run_manifest = {
        "run_id": run_id,
        "created_at_utc": datetime.now(timezone.utc).isoformat(),
        "git_commit": _git_commit(repo),
        "dataset_archive_sha256": "ee88e12be5acd05d0c1ec08bcebdfb4f72f15789c91dfddfc4f96277826adbf3",
        "source_workbook_sha256": processing["source_workbook_sha256"],
        "processed_data_sha256": processing["processed_sha256"],
        "config_sha256": _sha(config_path),
        "split_version": split["version"],
        "seeds": cfg["seeds"],
        "sequence_windows": cfg["sequence_windows"],
        "replay_command": f"PYTHONPATH=01_Product_Source/GFIS_Project python3 01_Product_Source/GFIS_Project/scripts/gate3_run_experiment.py --run-id {run_id}-replay",
        "artifacts": artifacts,
    }
    (run_dir / "run_manifest.json").write_text(json.dumps(run_manifest, indent=2) + "\n")
    return run_dir
