#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Download Copernicus Marine SEALEVEL_GLO_PHY_L4_MY_008_047
(cmems_obs-sl_glo_phy-ssh_my_allsat-demo-l4-duacs-0.125deg_P1D-i)
year by year (1993-2024) into per-year folders, via direct zarr read.

Variables (8):
  adt, sla, tpa_correction, err_sla, ugos, ugosa, vgos, vgosa

Only operates under /data/public/lzy/test.
"""
import os
import sys
import time
import traceback
import xarray as xr

# ---------- config ----------
ZARR_URL = (
    "https://s3.waw3-1.cloudferro.com/mdl-arco-time-045/arco/"
    "SEALEVEL_GLO_PHY_L4_MY_008_047/"
    "cmems_obs-sl_glo_phy-ssh_my_allsat-demo-l4-duacs-0.125deg_P1D-i_202511/"
    "timeChunked.zarr"
)
BASE_DIR = "/data/public/lzy/test"
DATA_ROOT = os.path.join(BASE_DIR, "sealevel_glo_phy_l4_my")

START_YEAR = 1993
END_YEAR = 2024

VARIABLES = [
    "adt", "sla", "tpa_correction", "err_sla",
    "ugos", "ugosa", "vgos", "vgosa",
]

# dataset ends 2024-11-19
END_DATE = "2024-11-19"

MAX_RETRIES = 3
RETRY_WAIT = 60


def year_completed(year_dir, year):
    """Check if a year folder already has a completed .nc file."""
    if not os.path.isdir(year_dir):
        return False
    expected = os.path.join(
        year_dir,
        f"cmems_obs-sl_glo_phy-ssh_my_allsat-demo-l4-duacs-0.125deg_P1D-i_{year}.nc",
    )
    if os.path.isfile(expected) and os.path.getsize(expected) > 1024 * 1024:
        return True
    return False


def check_disk(path):
    """Return available GB."""
    st = os.statvfs(path)
    return st.f_bavail * st.f_frsize / (1024 ** 3)


def download_year(ds, year):
    year_dir = os.path.join(DATA_ROOT, str(year))
    os.makedirs(year_dir, exist_ok=True)

    if year_completed(year_dir, year):
        print(f"[{year}] already downloaded, skipping.", flush=True)
        return True

    avail_gb = check_disk(DATA_ROOT)
    print(f"[{year}] disk available: {avail_gb:.1f} GB", flush=True)
    if avail_gb < 5:
        print(f"[{year}] WARNING: low disk space ({avail_gb:.1f} GB)", flush=True)

    start = f"{year}-01-01"
    if year == 2024:
        end = END_DATE
    else:
        end = f"{year}-12-31"

    out_name = f"cmems_obs-sl_glo_phy-ssh_my_allsat-demo-l4-duacs-0.125deg_P1D-i_{year}.nc"
    out_path = os.path.join(year_dir, out_name)

    print(f"[{year}] extracting {start} -> {end}, vars={VARIABLES}", flush=True)

    for attempt in range(1, MAX_RETRIES + 1):
        try:
            t0 = time.time()
            sub = ds[VARIABLES].sel(time=slice(start, end))
            # write to temp then rename for atomicity
            tmp_path = out_path + ".part"
            sub.to_netcdf(tmp_path)
            os.rename(tmp_path, out_path)
            size_gb = os.path.getsize(out_path) / (1024 ** 3)
            print(
                f"[{year}] done: {size_gb:.2f} GB in {time.time()-t0:.0f}s",
                flush=True,
            )
            return True
        except Exception as e:
            print(f"[{year}] attempt {attempt}/{MAX_RETRIES} failed: {e}", flush=True)
            traceback.print_exc()
            # clean partial
            tmp_path = out_path + ".part"
            if os.path.exists(tmp_path):
                os.remove(tmp_path)
            if attempt < MAX_RETRIES:
                time.sleep(RETRY_WAIT)
    print(f"[{year}] FAILED after {MAX_RETRIES} attempts.", flush=True)
    return False


def main():
    os.makedirs(DATA_ROOT, exist_ok=True)
    print(f"Base dir: {DATA_ROOT}", flush=True)
    print(f"Years: {START_YEAR}-{END_YEAR}", flush=True)
    print(f"Variables: {VARIABLES}", flush=True)
    print(f"Zarr: {ZARR_URL}", flush=True)

    # optional: single year via argv
    only = None
    if len(sys.argv) > 1:
        only = int(sys.argv[1])

    print("Opening zarr store...", flush=True)
    t0 = time.time()
    ds = xr.open_zarr(ZARR_URL, consolidated=True, chunks={})
    print(f"Opened in {time.time()-t0:.1f}s", flush=True)
    print(f"time range: {ds.time.values[0]} -> {ds.time.values[-1]}", flush=True)

    failed = []
    for year in range(START_YEAR, END_YEAR + 1):
        if only is not None and year != only:
            continue
        ok = download_year(ds, year)
        if not ok:
            failed.append(year)

    print("\n===== SUMMARY =====", flush=True)
    if failed:
        print("Failed years:", failed, flush=True)
    else:
        print("All years completed.", flush=True)
    return 0 if not failed else 1


if __name__ == "__main__":
    sys.exit(main())