#!/usr/bin/env python3
import argparse
from pathlib import Path
import traceback

import xarray as xr


def get_args():
    p = argparse.ArgumentParser()
    p.add_argument("--model", required=True)
    p.add_argument("--outdir", required=True)
    p.add_argument(
        "--vars",
        default="sst",
        help="Comma-separated variable list, e.g. sst or sst,prec,tref,prmsl",
    )
    p.add_argument(
        "--s-chunk",
        type=int,
        default=12,
        help="Number of forecast-start months per download chunk",
    )
    return p.parse_args()


def parse_vars(var_string):
    return [v.strip() for v in var_string.split(",") if v.strip()]


def get_candidate_urls(base_url, model, varname):
    """
    IRI NMME usually exposes variables as separate datasets, e.g.
      .MONTHLY/.sst/dods
      .MONTHLY/.prec/dods
      .MONTHLY/.tref/dods

    Therefore URL search is done per model and per variable.
    """
    return [
        # Preferred hindcast-style paths
        f"{base_url}/.{model}/.HINDCAST/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.HINDCAST/.{varname}/dods",

        # Nested CanSIPS-style component systems
        f"{base_url}/.{model}/.CanCM4i-IC3/.HINDCAST/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.GEM5-NEMO/.HINDCAST/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.CanESM5/.HINDCAST/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.GEM5.2-NEMO/.HINDCAST/.MONTHLY/.{varname}/dods",

        # Legacy monthly archive structures
        f"{base_url}/.{model}/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.mc8210/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.sc8210/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.mc8110/.MONTHLY/.{varname}/dods",
        f"{base_url}/.{model}/.sc8110/.MONTHLY/.{varname}/dods",
    ]


def open_first_valid_url(model, varname):
    base_url = "https://iridl.ldeo.columbia.edu/SOURCES/.Models/.NMME"

    for url in get_candidate_urls(base_url, model, varname):
        try:
            ds = xr.open_dataset(url, decode_times=False)

            if varname in ds.data_vars:
                _ = dict(ds.sizes)
                return url, ds

            ds.close()
        except Exception:
            continue

    return None, None


def safe_coord_value(ds, dim, idx):
    try:
        if dim in ds.coords:
            return str(ds[dim].isel({dim: idx}).values)
    except Exception:
        pass
    return str(idx)


def make_s_chunks(n_s, s_chunk):
    return [(i, min(i + s_chunk, n_s)) for i in range(0, n_s, s_chunk)]


def make_encoding(ds_sub, varname, s_chunk):
    enc = {}

    if varname not in ds_sub:
        return enc

    chunks = []
    for d in ds_sub[varname].dims:
        n = ds_sub.sizes[d]

        if d == "S":
            chunks.append(min(n, s_chunk))
        elif d == "M":
            chunks.append(1)
        elif d == "L":
            chunks.append(min(n, 12))
        elif d in ("P", "P2"):
            chunks.append(min(n, 2))
        elif d == "Y":
            chunks.append(min(n, 45))
        elif d == "X":
            chunks.append(min(n, 72))
        else:
            chunks.append(n)

    enc[varname] = {
        "zlib": True,
        "complevel": 1,
        "chunksizes": tuple(chunks),
    }

    return enc


def download_variable_for_model(model, varname, outdir, s_chunk, log):
    log("")
    log("-" * 72)
    log(f"Variable: {varname}")

    url, ds = open_first_valid_url(model, varname)

    if url is None:
        log(f"SKIP: no usable URL found for {model} variable {varname}")
        return {
            "varname": varname,
            "status": "SKIP",
            "url": "",
            "message": "No usable URL found",
        }

    log(f"URL: {url}")
    log(f"sizes: {dict(ds.sizes)}")

    if "S" not in ds.sizes:
        log(f"SKIP: {model} {varname} has no S dimension")
        ds.close()
        return {
            "varname": varname,
            "status": "SKIP",
            "url": url,
            "message": "No S dimension",
        }

    n_s = ds.sizes["S"]
    s_chunks = make_s_chunks(n_s, s_chunk)

    if "M" in ds.sizes:
        m_indices = list(range(ds.sizes["M"]))
    else:
        m_indices = [None]

    var_dir = outdir / model / varname
    var_dir.mkdir(parents=True, exist_ok=True)

    with open(var_dir / "source_info.txt", "w") as f:
        f.write(f"model: {model}\n")
        f.write(f"variable: {varname}\n")
        f.write(f"url: {url}\n")
        f.write(f"sizes: {dict(ds.sizes)}\n")
        f.write(f"S_start: {safe_coord_value(ds, 'S', 0)}\n")
        f.write(f"S_end: {safe_coord_value(ds, 'S', n_s - 1)}\n")
        if "M" in ds.sizes:
            f.write(f"M_start: {safe_coord_value(ds, 'M', 0)}\n")
            f.write(f"M_end: {safe_coord_value(ds, 'M', ds.sizes['M'] - 1)}\n")
        if "L" in ds.sizes:
            f.write(f"L_start: {safe_coord_value(ds, 'L', 0)}\n")
            f.write(f"L_end: {safe_coord_value(ds, 'L', ds.sizes['L'] - 1)}\n")

    for m_idx in m_indices:
        if m_idx is None:
            member_label = "M000"
            ds_m = ds
        else:
            member_label = f"M{m_idx + 1:03d}"
            ds_m = ds.isel(M=m_idx)

        member_dir = var_dir / member_label
        member_dir.mkdir(parents=True, exist_ok=True)

        log("")
        log(f"Downloading {model} {varname} {member_label}")

        chunk_files = []

        for s0, s1 in s_chunks:
            chunk_file = member_dir / f"{varname}_{model}_{member_label}_S{s0:04d}-{s1 - 1:04d}.nc"
            tmp_file = member_dir / f".tmp_{chunk_file.name}"

            if chunk_file.exists() and chunk_file.stat().st_size > 0:
                log(f"  exists: {chunk_file.name}")
                chunk_files.append(chunk_file)
                continue

            s_start_label = safe_coord_value(ds, "S", s0)
            s_end_label = safe_coord_value(ds, "S", s1 - 1)

            log(f"  S index {s0}:{s1 - 1} ({s_start_label} to {s_end_label})")

            try:
                ds_sub = ds_m[[varname]].isel(S=slice(s0, s1)).load()
                enc = make_encoding(ds_sub, varname, s_chunk)

                if tmp_file.exists():
                    tmp_file.unlink()

                ds_sub.to_netcdf(
                    tmp_file,
                    format="NETCDF4_CLASSIC",
                    engine="netcdf4",
                    encoding=enc,
                )

                tmp_file.rename(chunk_file)
                chunk_files.append(chunk_file)

            except Exception as exc:
                log(f"  WARNING: failed {model} {varname} {member_label} S={s0}:{s1 - 1}")
                log(f"  {exc}")
                log(traceback.format_exc())

                if tmp_file.exists():
                    tmp_file.unlink()

                continue

        with open(member_dir / "chunk_files.txt", "w") as f:
            for cf in chunk_files:
                f.write(str(cf) + "\n")

    ds.close()

    log(f"DONE variable: {varname}")

    return {
        "varname": varname,
        "status": "DONE",
        "url": url,
        "message": "done",
    }


def main():
    args = get_args()

    model = args.model
    outdir = Path(args.outdir)
    s_chunk = args.s_chunk
    requested_vars = parse_vars(args.vars)

    model_dir = outdir / model
    model_dir.mkdir(parents=True, exist_ok=True)

    log_file = model_dir / "download.log"
    if log_file.exists():
        log_file.unlink()

    def log(msg):
        print(msg, flush=True)
        with open(log_file, "a") as f:
            f.write(msg + "\n")

    log("")
    log("=" * 72)
    log(f"Processing model: {model}")
    log(f"Requested variables: {requested_vars}")

    records = []

    for varname in requested_vars:
        rec = download_variable_for_model(
            model=model,
            varname=varname,
            outdir=outdir,
            s_chunk=s_chunk,
            log=log,
        )
        records.append(rec)

    done_vars = [r["varname"] for r in records if r["status"] == "DONE"]
    skipped_vars = [r["varname"] for r in records if r["status"] != "DONE"]

    with open(model_dir / "model_summary.txt", "w") as f:
        f.write(f"model: {model}\n")
        f.write(f"requested_vars: {requested_vars}\n")
        f.write(f"done_vars: {done_vars}\n")
        f.write(f"skipped_vars: {skipped_vars}\n")
        f.write("\nrecords:\n")
        for r in records:
            f.write(str(r) + "\n")

    if done_vars:
        (model_dir / "DONE").write_text("done\n")
        log("")
        log(f"DONE model: {model}")
        log(f"Downloaded variables: {done_vars}")
        if skipped_vars:
            log(f"Skipped variables: {skipped_vars}")
        return 0

    (model_dir / "FAILED").write_text("No requested variables downloaded\n")
    log("")
    log(f"FAILED model: {model}; no requested variables downloaded")
    return 2


if __name__ == "__main__":
    raise SystemExit(main())
