#!/usr/bin/env python3
"""Summarize ELM history coverage and selected ILAMB reference products."""

from __future__ import annotations

import argparse
import glob
import os

import netCDF4


DEFAULT_VARIABLES = (
    "H2OSNO",
    "SNOWDP",
    "SNOW_DEPTH",
    "FSNO",
    "FSNO_EFF",
    "QSNOMELT",
    "SNOW",
    "RAIN",
    "QRUNOFF",
    "QOVER",
    "QDRAI",
    "SOILWATER_10CM",
    "H2OSOI",
    "TWS",
    "QVEGE",
    "QVEGT",
    "QSOIL",
    "EFLX_LH_TOT",
    "TSOI",
    "ALT",
    "ALTMAX",
    "FSR",
    "FSDS",
    "TSA",
)


def describe_time(dataset: netCDF4.Dataset) -> str:
    if "time" not in dataset.variables:
        return "no time coordinate"
    time = dataset.variables["time"]
    if time.size == 0:
        return "empty time coordinate"
    units = getattr(time, "units", None)
    calendar = getattr(time, "calendar", "standard")
    if not units:
        return f"time={time[0]}..{time[-1]} (units missing)"
    dates = netCDF4.num2date([time[0], time[-1]], units, calendar=calendar)
    return f"{dates[0]} .. {dates[-1]} ({calendar})"


def summarize_models(paths: list[str], variables: tuple[str, ...]) -> None:
    for path in paths:
        files = sorted(glob.glob(os.path.join(path, "*.elm.h0.*.nc")))
        print(f"MODEL {path}")
        print(f"  files: {len(files)}")
        if not files:
            continue
        print(f"  first: {os.path.basename(files[0])}")
        print(f"  last:  {os.path.basename(files[-1])}")
        with netCDF4.Dataset(files[0]) as dataset:
            print(f"  first-file time: {describe_time(dataset)}")
            for name in variables:
                if name not in dataset.variables:
                    continue
                var = dataset.variables[name]
                print(
                    f"  {name}: dims={var.dimensions}, "
                    f"units={getattr(var, 'units', '')!r}, "
                    f"long_name={getattr(var, 'long_name', '')!r}"
                )
        with netCDF4.Dataset(files[-1]) as dataset:
            print(f"  last-file time:  {describe_time(dataset)}")


def summarize_references(root: str) -> None:
    variables = (
        "active_layer_thickness",
        "albedo",
        "evspsbl",
        "hfls",
        "mrro",
        "mrso",
        "mrsol",
        "mrsos",
        "permafrost",
        "pr",
        "snw",
        "swe",
        "tran",
        "tws",
        "twsa",
    )
    data_root = os.path.join(root, "DATA")
    for group in variables:
        for path in sorted(glob.glob(os.path.join(data_root, group, "**", "*.nc"), recursive=True)):
            try:
                with netCDF4.Dataset(path) as dataset:
                    candidates = []
                    for name, var in dataset.variables.items():
                        if name.lower() in {
                            "time",
                            "time_bnds",
                            "time_bounds",
                            "lat",
                            "latitude",
                            "lon",
                            "longitude",
                            "area",
                            "landfrac",
                            "data_bnds",
                            "nb",
                        }:
                            continue
                        if len(var.dimensions) > 0:
                            candidates.append(
                                f"{name}[{getattr(var, 'units', '')}]"
                            )
                    relpath = os.path.relpath(path, root)
                    print(
                        f"REFERENCE {relpath}: {describe_time(dataset)}; "
                        f"variables={','.join(candidates)}"
                    )
            except Exception as exc:  # keep inventorying if one optional file is bad
                print(f"REFERENCE {path}: ERROR {exc}")


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--ilamb-root", required=True)
    parser.add_argument("model_paths", nargs="+")
    args = parser.parse_args()
    summarize_models(args.model_paths, DEFAULT_VARIABLES)
    summarize_references(args.ilamb_root)


if __name__ == "__main__":
    main()
