#!/usr/bin/env python3
"""Create and optionally submit the Pathfinder ILAMB hydrology/snow study.

The generated study compares the standard ELM run and shallow-snow ELM run
over model years 1980--2022. ILAMB automatically restricts each confrontation
to the overlap with its reference product inside that study window.
"""

from __future__ import annotations

import argparse
import datetime as dt
import hashlib
import json
import os
import platform
import re
import subprocess
from pathlib import Path

import netCDF4


PROJECT_ROOT = Path("/projects/hpcl-cli185/proj-shared/zdr")
DEFAULT_CONTROL = PROJECT_ROOT / "20260625.ERA5r025_r05_IcoswISC30E3r5.pm-cpu.20TR"
DEFAULT_SHALLOW = (
    PROJECT_ROOT
    / "20260915.ERA5r025_r05_IcoswISC30E3r5.pm-cpu.20TR.shallowSnow"
)
DEFAULT_WORKDIR = PROJECT_ROOT / "ilamb_runs/hydrology_snow_1980_2022"
DEFAULT_ILAMB_ROOT = PROJECT_ROOT / "ILAMB_DATA"
DEFAULT_ENV_SCRIPT = PROJECT_ROOT / "load_ilamb_pf_sysmpi.sh"
REGIONS = (
    "global",
    "nwn",  # IPCC AR6 NWN: Northwest North America
    "nec",  # IPCC AR6 NEC: Northeast Canada
    "neu",  # IPCC AR6 NEU: Northern Europe
    "wsb",  # IPCC AR6 WSB: Western Siberia
    "esb",  # IPCC AR6 ESB: Eastern Siberia
    "rfe",  # IPCC AR6 RFE: Russian Far East
    "rar",  # IPCC AR6 RAR: Russian Arctic
)
REQUIRED_MODEL_VARIABLES = {
    "H2OSNO",
    "QRUNOFF",
    "SOILWATER_10CM",
    "TWS",
    "QVEGE",
    "QVEGT",
    "QSOIL",
    "RAIN",
    "SNOW",
    "TSOI",
    "ALTMAX",
}
REFERENCE_FILES = (
    "DATA/regions/GlobalLand.nc",
    "DATA/regions/IPCCRegions.nc",
    "DATA/swe/CanSISE/swe.nc",
    "DATA/snw/CARDAMOM/snw.nc",
    "DATA/evspsbl/GLEAMv3.3a/et.nc",
    "DATA/evspsbl/MODIS/et_0.5x0.5.nc",
    "DATA/evspsbl/MOD16A2/et.nc",
    "DATA/mrro/Dai/runoff.nc",
    "DATA/mrro/LORA/LORA.nc",
    "DATA/mrro/CLASS/mrro.nc",
    "DATA/twsa/GRACE/twsa_0.5x0.5.nc",
    "DATA/mrsos/WangMao/mrsos_olc.nc",
    "DATA/permafrost/Brown2002/Brown2002.nc",
    "DATA/permafrost/Obu2018/Obu2018.nc",
    "DATA/active_layer_thickness/CALM/CALM.nc",
    "DATA/pr/GPCCv2018/pr.nc",
    "DATA/pr/GPCPv2.3/pr.nc",
)


CONFIG_TEXT = """\
# Reproducible ILAMB 2.7 configuration for ELM hydrology and snow.
# The variable mappings follow ILAMB's installed ilamb_nohoff_final_CLM.cfg.
#! define_regions = DATA/regions/GlobalLand.nc,DATA/regions/IPCCRegions.nc

[h1: Snow and Frozen Ground]
bgcolor = "#DDEEFF"

[h2: Snow Water Equivalent]
variable       = "H2OSNO"
alternate_vars = "swe,snw"
cmap           = "Blues"
weight         = 10
ctype          = "ConfSWE"

[CanSISE]
source     = "DATA/swe/CanSISE/swe.nc"
weight     = 25
plot_unit  = "cm"
table_unit = "cm"

[CARDAMOM]
source     = "DATA/snw/CARDAMOM/snw.nc"
weight     = 25
plot_unit  = "cm"
table_unit = "cm"

[h2: Active Layer Thickness]
variable       = "ALTMAX"
alternate_vars = "alt"
cmap           = "Blues"
weight         = 5

[CALM]
source     = "DATA/active_layer_thickness/CALM/CALM.nc"
weight     = 25
plot_unit  = "m"
table_unit = "m"

[h2: Permafrost]
variable       = "TSOI"
alternate_vars = "tsl"
weight         = 3

[Brown2002]
ctype  = "ConfPermafrost"
source = "DATA/permafrost/Brown2002/Brown2002.nc"
y0     = 1985.
yf     = 2005.
Teps   = 273.15
dmax   = 3.5

[Obu2018]
ctype  = "ConfPermafrost"
source = "DATA/permafrost/Obu2018/Obu2018.nc"
y0     = 2000.
yf     = 2016.
Teps   = 273.15
dmax   = 3.5

[h1: Hydrology Cycle]
bgcolor = "#E6F9FF"

[h2: Evapotranspiration]
variable       = "et"
alternate_vars = "evspsbl"
derived        = "QVEGE+QVEGT+QSOIL"
cmap           = "Blues"
weight         = 5
mass_weighting = True

[GLEAMv3.3a]
source     = "DATA/evspsbl/GLEAMv3.3a/et.nc"
weight     = 15
table_unit = "mm d-1"
plot_unit  = "mm d-1"

[MODIS]
source     = "DATA/evspsbl/MODIS/et_0.5x0.5.nc"
weight     = 15
table_unit = "mm d-1"
plot_unit  = "mm d-1"

[MOD16A2]
source     = "DATA/evspsbl/MOD16A2/et.nc"
weight     = 15
table_unit = "mm d-1"
plot_unit  = "mm d-1"

[h2: Runoff]
variable       = "runoff"
alternate_vars = "mrro,QRUNOFF"
weight         = 5
mass_weighting = True

[Dai]
ctype  = "ConfRunoff"
source = "DATA/mrro/Dai/runoff.nc"
weight = 15

[LORA]
source     = "DATA/mrro/LORA/LORA.nc"
table_unit = "mm d-1"
plot_unit  = "mm d-1"
weight     = 15

[CLASS]
source     = "DATA/mrro/CLASS/mrro.nc"
plot_unit  = "mm d-1"
table_unit = "mm d-1"
weight     = 25

[h2: Terrestrial Water Storage Anomaly]
variable       = "twsa"
alternate_vars = "tws,TWS"
derived        = "RAIN+SNOW-QVEGE-QVEGT-QSOIL-QRUNOFF"
cmap           = "Blues"
weight         = 5
ctype          = "ConfTWSA"

[GRACE]
source = "DATA/twsa/GRACE/twsa_0.5x0.5.nc"
weight = 25

[h2: Surface Soil Moisture]
variable       = "SOILWATER_10CM"
alternate_vars = "mrsos"
weight         = 3
cmap           = "Blues"

[WangMao]
source = "DATA/mrsos/WangMao/mrsos_olc.nc"
weight = 15

[h2: Precipitation]
variable       = "pr"
derived        = "RAIN+SNOW"
cmap           = "Blues"
weight         = 2
mass_weighting = True

[GPCCv2018]
source     = "DATA/pr/GPCCv2018/pr.nc"
land       = True
weight     = 20
table_unit = "mm d-1"
plot_unit  = "mm d-1"
space_mean = True

[GPCPv2.3]
source     = "DATA/pr/GPCPv2.3/pr.nc"
land       = True
weight     = 20
table_unit = "mm d-1"
plot_unit  = "mm d-1"
space_mean = True
"""


def sha256(text: str) -> str:
    return hashlib.sha256(text.encode("utf-8")).hexdigest()


def atomic_write(path: Path, text: str) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary = path.with_suffix(path.suffix + ".tmp")
    temporary.write_text(text, encoding="utf-8")
    temporary.replace(path)


def expected_months(start_year: int, end_year: int) -> list[str]:
    return [
        f"{year:04d}-{month:02d}"
        for year in range(start_year, end_year + 1)
        for month in range(1, 13)
    ]


def validate_model(path: Path, start_year: int, end_year: int) -> dict:
    if not path.is_dir():
        raise FileNotFoundError(f"Model directory does not exist: {path}")
    pattern = re.compile(r"\.elm\.h0\.(\d{4}-\d{2})\.nc$")
    by_month: dict[str, Path] = {}
    for filename in sorted(path.glob("*.elm.h0.*.nc")):
        match = pattern.search(filename.name)
        if match:
            by_month[match.group(1)] = filename
    wanted = expected_months(start_year, end_year)
    missing = [month for month in wanted if month not in by_month]
    if missing:
        raise RuntimeError(
            f"{path} is missing {len(missing)} required monthly files; "
            f"first missing months: {missing[:12]}"
        )
    selected = [by_month[month] for month in wanted]
    with netCDF4.Dataset(selected[0]) as dataset:
        missing_variables = sorted(REQUIRED_MODEL_VARIABLES - set(dataset.variables))
        if missing_variables:
            raise RuntimeError(f"{selected[0]} lacks variables {missing_variables}")
        grid = {
            "lat": len(dataset.dimensions["lat"]),
            "lon": len(dataset.dimensions["lon"]),
        }
    return {
        "path": str(path),
        "file_count": len(selected),
        "first_file": str(selected[0]),
        "last_file": str(selected[-1]),
        "grid": grid,
    }


def validate_references(ilamb_root: Path) -> None:
    missing = [str(ilamb_root / relpath) for relpath in REFERENCE_FILES if not (ilamb_root / relpath).is_file()]
    if missing:
        raise FileNotFoundError("Missing ILAMB reference files:\n" + "\n".join(missing))


def make_model_setup(control: Path, shallow: Path) -> str:
    return (
        "# Model name, absolute directory, group\n"
        f"ELM_Control, {control}, ELM snow experiment\n"
        f"ELM_ShallowSnow, {shallow}, ELM snow experiment\n"
    )


def make_sbatch(
    workdir: Path,
    ilamb_root: Path,
    env_script: Path,
    start_year: int,
    end_year: int,
    ranks: int,
    memory: str,
    walltime: str,
    clean: bool,
) -> str:
    clean_option = " --clean" if clean else ""
    region_options = " ".join(REGIONS)
    return f"""\
#!/bin/bash
#SBATCH -A hpcl-cli185
#SBATCH -p serial
#SBATCH -q normal
#SBATCH -N 1
#SBATCH -n {ranks}
#SBATCH -c 1
#SBATCH --mem={memory}
#SBATCH -t {walltime}
#SBATCH -J ilamb_hydro_snow
#SBATCH -o {workdir}/logs/ilamb_hydro_snow_%j.out
#SBATCH -e {workdir}/logs/ilamb_hydro_snow_%j.err

set -eo pipefail
umask 002

set +u
source /etc/profile || true
source {env_script}
set -u

export ILAMB_ROOT={ilamb_root}
export MPLBACKEND=Agg
export OMP_NUM_THREADS=1
export PRTE_MCA_prte_tmpdir_base=/tmp

WORKDIR={workdir}
export PYTHONPATH="$WORKDIR/runtime_patch:${{PYTHONPATH:-}}"
mkdir -p "$WORKDIR/logs" "$WORKDIR/build"
cd "$WORKDIR"

echo "hostname=$(hostname)"
echo "date_start=$(date --iso-8601=seconds)"
echo "ilamb=$(command -v ilamb-run)"
echo "ILAMB_ROOT=$ILAMB_ROOT"
echo "ranks=$SLURM_NTASKS"

mpiexec -n "$SLURM_NTASKS" ilamb-run \\
  --config "$WORKDIR/ilamb_hydrology_snow.cfg" \\
  --model_setup "$WORKDIR/models.txt" \\
  --filter ".elm.h0." \\
  --study_limits {start_year} {end_year} \\
  --regions {region_options} \\
  --build_dir "$WORKDIR/build" \\
  --title "ELM Control vs Shallow Snow: Hydrology and Snow ({start_year}-{end_year})"{clean_option}

status=$?
echo "date_end=$(date --iso-8601=seconds)"
echo "status=$status"
exit "$status"
"""


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--control", type=Path, default=DEFAULT_CONTROL)
    parser.add_argument("--shallow-snow", type=Path, default=DEFAULT_SHALLOW)
    parser.add_argument("--workdir", type=Path, default=DEFAULT_WORKDIR)
    parser.add_argument("--ilamb-root", type=Path, default=DEFAULT_ILAMB_ROOT)
    parser.add_argument("--env-script", type=Path, default=DEFAULT_ENV_SCRIPT)
    parser.add_argument("--start-year", type=int, default=1980)
    parser.add_argument("--end-year", type=int, default=2022)
    parser.add_argument("--ranks", type=int, default=8)
    parser.add_argument("--memory", default="128G")
    parser.add_argument("--walltime", default="12:00:00")
    parser.add_argument("--clean", action="store_true", help="Force ILAMB to recompute cached confrontations")
    parser.add_argument("--submit", action="store_true", help="Submit the generated SLURM job")
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    if args.start_year > args.end_year:
        raise ValueError("start-year must not exceed end-year")
    if args.ranks < 1:
        raise ValueError("ranks must be positive")
    if not args.env_script.is_file():
        raise FileNotFoundError(args.env_script)
    patch_path = Path(__file__).resolve().parent / "runtime_patch/sitecustomize.py"
    if not patch_path.is_file():
        raise FileNotFoundError(patch_path)

    models = [
        validate_model(args.control, args.start_year, args.end_year),
        validate_model(args.shallow_snow, args.start_year, args.end_year),
    ]
    if models[0]["grid"] != models[1]["grid"]:
        raise RuntimeError(f"Model grids differ: {models[0]['grid']} vs {models[1]['grid']}")
    validate_references(args.ilamb_root)

    args.workdir.mkdir(parents=True, exist_ok=True)
    (args.workdir / "logs").mkdir(exist_ok=True)
    config_path = args.workdir / "ilamb_hydrology_snow.cfg"
    models_path = args.workdir / "models.txt"
    sbatch_path = args.workdir / "run_ilamb_hydrology_snow.sbatch"
    model_text = make_model_setup(args.control, args.shallow_snow)
    sbatch_text = make_sbatch(
        args.workdir,
        args.ilamb_root,
        args.env_script,
        args.start_year,
        args.end_year,
        args.ranks,
        args.memory,
        args.walltime,
        args.clean,
    )
    atomic_write(config_path, CONFIG_TEXT)
    atomic_write(models_path, model_text)
    atomic_write(sbatch_path, sbatch_text)
    sbatch_path.chmod(0o750)

    manifest = {
        "created_utc": dt.datetime.now(dt.timezone.utc).isoformat(),
        "created_on": platform.node(),
        "study_years_inclusive": [args.start_year, args.end_year],
        "regions": list(REGIONS),
        "models": models,
        "ilamb_root": str(args.ilamb_root),
        "environment_script": str(args.env_script),
        "slurm": {
            "ranks": args.ranks,
            "memory": args.memory,
            "walltime": args.walltime,
            "clean": args.clean,
        },
        "files": {
            config_path.name: sha256(CONFIG_TEXT),
            models_path.name: sha256(model_text),
            sbatch_path.name: sha256(sbatch_text),
            "runtime_patch/sitecustomize.py": hashlib.sha256(
                patch_path.read_bytes()
            ).hexdigest(),
        },
        "reference_files": list(REFERENCE_FILES),
    }
    manifest_path = args.workdir / "manifest.json"
    atomic_write(manifest_path, json.dumps(manifest, indent=2) + "\n")

    print(f"Validated {sum(model['file_count'] for model in models)} model files")
    print(f"Wrote workflow inputs to {args.workdir}")
    if args.submit:
        result = subprocess.run(
            ["sbatch", "--parsable", str(sbatch_path)],
            check=True,
            text=True,
            capture_output=True,
        )
        job_id = result.stdout.strip().split(";")[0]
        atomic_write(args.workdir / "job_id.txt", job_id + "\n")
        print(f"Submitted SLURM job {job_id}")


if __name__ == "__main__":
    main()
