#!/usr/bin/env python3
"""
HYCOM GLBv0.08 Reanalysis (expt_53.X) Downloader
Dataset  : Global 1/12° Reanalysis, 1994-01-01 – 2015-12-30
Domain   : 10–32°N, 100–126°E  (±1° buffer for ROMS boundaries)
Variables: water_temp, salinity, water_u, water_v, surf_el
Temporal : ~3-hourly (~7.9 steps/day)
Output   : ./hycom_data/hycom_YYYYMM.nc  (one file per month)

Strategy : OPeNDAP requests are capped at CHUNK_DAYS days to avoid DAP
           server-side timeouts. Chunks are retried up to MAX_RETRY times,
           then merged in memory and saved as a monthly NetCDF.

Dependencies:
    pip install xarray netCDF4 numpy pandas
"""

import sys
import time
import argparse
import logging
import xarray as xr
import pandas as pd
import numpy as np
from pathlib import Path

# ─── Configuration ────────────────────────────────────────────────────────────
LAT_S, LAT_N = 9.0, 33.0      # ±1° buffer beyond 10–32°N
LON_W, LON_E = 99.0, 127.0    # ±1° buffer beyond 100–126°E
DATE_START   = '1997-01-01'
DATE_END     = '1999-12-31'
OUTPUT_DIR   = Path('./hycom_data')

# Per-year OPeNDAP aggregation (faster than the 22-year aggregate)
OPENDAP_TMPL = 'https://tds.hycom.org/thredds/dodsC/GLBv0.08/expt_53.X/data/{year}'

VARS_3D  = ['water_temp', 'salinity', 'water_u', 'water_v']
VARS_2D  = ['surf_el']
ALL_VARS = VARS_3D + VARS_2D

CHUNK_DAYS = 5      # days per OPeNDAP request — tune down if DAP failures persist
MAX_RETRY  = 3      # retries per chunk
RETRY_WAIT = 10     # seconds between retries

COMPRESS = {'zlib': True, 'complevel': 4}

# ─── Logging ──────────────────────────────────────────────────────────────────
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s [%(levelname)s] %(message)s',
    handlers=[logging.StreamHandler(sys.stdout)],
)
log = logging.getLogger(__name__)


def open_year(year: int):
    url = OPENDAP_TMPL.format(year=year)
    log.info(f'Opening {year}: {url}')
    try:
        ds = xr.open_dataset(url, engine='netcdf4', drop_variables=['tau'])
        n  = ds.sizes['time']
        log.info(f'  {year}: {n} time steps (~{n/365:.1f}/day), '
                 f'vars={[v for v in ALL_VARS if v in ds.data_vars]}')
        return ds
    except Exception as exc:
        log.error(f'Cannot open {year}: {exc}')
        return None


def fetch_chunk(ds, avail, t_start: str, t_end: str):
    """Fetch one time chunk with retries. Returns xr.Dataset or None."""
    for attempt in range(1, MAX_RETRY + 1):
        try:
            sub = ds[avail].sel(
                time=slice(t_start, t_end),
                lat=slice(LAT_S, LAT_N),
                lon=slice(LON_W, LON_E),
            )
            sub.load()   # trigger actual OPeNDAP transfer
            return sub
        except Exception as exc:
            log.warning(f'    Attempt {attempt}/{MAX_RETRY} failed '
                        f'({t_start}–{t_end}): {exc}')
            if attempt < MAX_RETRY:
                log.info(f'    Retrying in {RETRY_WAIT}s ...')
                time.sleep(RETRY_WAIT)
    log.error(f'    All {MAX_RETRY} attempts failed for chunk {t_start}–{t_end}')
    return None


def download_month(ds, year, month):
    month_start = pd.Timestamp(year=year, month=month, day=1)
    month_end   = month_start + pd.offsets.MonthEnd(0)
    tag         = month_start.strftime('%Y%m')
    out_file    = OUTPUT_DIR / f'hycom_{tag}.nc'

    if out_file.exists():
        log.info(f'[SKIP] {tag}  ({out_file.stat().st_size / 1e6:.1f} MB)')
        return True

    log.info(f'[DOWNLOAD] {tag} (chunk_days={CHUNK_DAYS}) ...')
    avail  = [v for v in ALL_VARS if v in ds.data_vars]

    # Build list of CHUNK_DAYS-long date windows covering the month
    chunks     = []
    chunk_start = month_start
    while chunk_start <= month_end:
        chunk_end = min(chunk_start + pd.Timedelta(days=CHUNK_DAYS - 1), month_end)
        chunks.append((str(chunk_start.date()), str(chunk_end.date())))
        chunk_start = chunk_end + pd.Timedelta(days=1)

    log.info(f'  {len(chunks)} chunks of up to {CHUNK_DAYS} days each')

    # Download each chunk
    parts = []
    for i, (t0, t1) in enumerate(chunks, 1):
        log.info(f'  chunk {i}/{len(chunks)}: {t0} – {t1}')
        part = fetch_chunk(ds, avail, t0, t1)
        if part is None:
            log.error(f'  Month {tag} aborted — chunk {t0}–{t1} failed.')
            return False
        parts.append(part)

    # Merge and save
    try:
        log.info(f'  Merging {len(parts)} chunks ...')
        merged = xr.concat(parts, dim='time')
        encoding = {v: COMPRESS for v in merged.data_vars}
        merged.to_netcdf(out_file, encoding=encoding)
        size_mb = out_file.stat().st_size / 1e6
        log.info(f'  Saved: {out_file}  ({size_mb:.1f} MB, '
                 f'{merged.sizes["time"]} time steps)')
        # Free memory
        for p in parts:
            p.close()
        merged.close()
        return True
    except Exception as exc:
        log.error(f'  Save failed {tag}: {exc}')
        if out_file.exists():
            out_file.unlink()
        return False


def verify_mode(chunk_days: int):
    """Download first CHUNK_DAYS days of Jan 1997 only."""
    global CHUNK_DAYS
    CHUNK_DAYS = chunk_days
    log.info(f'=== VERIFY: Jan 1997, first {CHUNK_DAYS} days ===')
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    ds = open_year(1997)
    if ds is None:
        sys.exit(1)

    avail = [v for v in ALL_VARS if v in ds.data_vars]
    t0    = '1997-01-01'
    t1    = f'1997-01-{CHUNK_DAYS:02d}'
    log.info(f'Fetching {t0} – {t1} ...')
    chunk = fetch_chunk(ds, avail, t0, t1)
    ds.close()

    if chunk is None:
        log.error('=== VERIFICATION FAILED ===')
        sys.exit(1)

    out_file = OUTPUT_DIR / f'hycom_verify_{CHUNK_DAYS}days.nc'
    chunk.to_netcdf(out_file, encoding={v: COMPRESS for v in chunk.data_vars})
    size_mb = out_file.stat().st_size / 1e6
    log.info(f'Saved: {out_file}  ({size_mb:.1f} MB, '
             f'{chunk.sizes["time"]} time steps)')
    chunk.close()
    log.info('=== VERIFICATION PASSED ===')


def main():
    global CHUNK_DAYS
    parser = argparse.ArgumentParser(description='Download HYCOM reanalysis data')
    parser.add_argument('--verify', action='store_true',
                        help='Quick test: download first CHUNK_DAYS of Jan 1997')
    parser.add_argument('--chunk-days', type=int, default=CHUNK_DAYS,
                        help=f'Days per OPeNDAP request (default: {CHUNK_DAYS})')
    args = parser.parse_args()
    CHUNK_DAYS = args.chunk_days

    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    fh = logging.FileHandler(OUTPUT_DIR / 'download_hycom.log')
    fh.setFormatter(logging.Formatter('%(asctime)s [%(levelname)s] %(message)s'))
    log.addHandler(fh)

    if args.verify:
        verify_mode(args.chunk_days)
        return

    months = pd.date_range(DATE_START, DATE_END, freq='MS')
    failed = []
    current_year, ds = None, None

    for m in months:
        if m.year != current_year:
            if ds is not None:
                ds.close()
            current_year = m.year
            ds = open_year(current_year)
            if ds is None:
                for skip_m in months[months.year == current_year]:
                    failed.append(skip_m.strftime('%Y%m'))
                continue

        if not download_month(ds, m.year, m.month):
            failed.append(m.strftime('%Y%m'))

    if ds is not None:
        ds.close()

    if failed:
        log.warning(f'Failed months ({len(failed)}): {failed}')
    else:
        log.info('All months downloaded successfully.')


if __name__ == '__main__':
    main()
