#!/usr/bin/env python3
"""
WOA23 Monthly Climatology Downloader
Source   : NOAA NCEI (no authentication required)
Variables: temperature (t), salinity (s)
Periods  : annual mean (00) + 12 monthly means (01–12)
Resolution: 1.00° (decav climatology)  — each file ~8–15 MB, 26 files total
            0.25° files are ~1 GB each (26 GB total); use --hires to enable.

Note: WOA23 provides *climatological* monthly means without year information.
      These serve as ROMS initial conditions and T/S relaxation fields.
      Interannual daily forcing comes from HYCOM and ERA5.

Output   : ./woa23_data/temperature/woa23_tMM.nc
           ./woa23_data/salinity/woa23_sMM.nc

Dependencies: standard library only
"""

import sys
import argparse
import logging
import urllib.request
import urllib.error
from pathlib import Path

# ─── Configuration ────────────────────────────────────────────────────────────
OUTPUT_DIR = Path('./woa23_data')

BASE_URL = 'https://www.ncei.noaa.gov/data/oceans/woa/WOA23/DATA'

WOA_VARS = {
    'temperature': 't',
    'salinity':    's',
}

# 00 = annual mean;  01–12 = monthly climatology
MONTHS = ['00'] + [f'{m:02d}' for m in range(1, 13)]

# Resolution options
GRID_1DEG  = ('1.00', '01')   # ~8–15 MB per file  (recommended)
GRID_025DEG = ('0.25', '04')  # ~1 GB  per file  (--hires flag)

CHUNK_SIZE = 1024 * 1024  # 1 MB read chunks for progress display

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


def build_url(var_long, var_short, month, grid, grid_code):
    return (
        f'{BASE_URL}/{var_long}/netcdf/decav/{grid}/'
        f'woa23_decav_{var_short}{month}_{grid_code}.nc'
    )


def fetch_with_progress(url: str, dest: Path) -> bool:
    """Download with chunked read, progress display, and size verification."""
    log.info(f'  GET {url}')
    try:
        req  = urllib.request.Request(url)
        resp = urllib.request.urlopen(req, timeout=120)
        content_length = resp.headers.get('Content-Length')
        expected = int(content_length) if content_length else None

        downloaded = 0
        with open(dest, 'wb') as f:
            while True:
                chunk = resp.read(CHUNK_SIZE)
                if not chunk:
                    break
                f.write(chunk)
                downloaded += len(chunk)
                if expected:
                    pct = downloaded / expected * 100
                    print(f'\r  {downloaded/1e6:.1f}/{expected/1e6:.1f} MB  '
                          f'({pct:.0f}%)', end='', flush=True)

        print()  # newline after progress
        actual = dest.stat().st_size

        if expected and actual < expected * 0.99:
            log.error(f'  Truncated: got {actual} bytes, expected {expected}')
            dest.unlink()
            return False

        log.info(f'  -> {dest.name}  ({actual / 1e6:.1f} MB)')
        return True

    except urllib.error.HTTPError as exc:
        log.warning(f'  -> HTTP {exc.code} {exc.reason}')
    except Exception as exc:
        log.warning(f'  -> {exc}')

    if dest.exists():
        dest.unlink()
    return False


def download_one(var_long, var_short, month, var_dir: Path,
                 grid, grid_code) -> bool:
    label = f'{var_long} month={month} ({grid}°)'
    dest  = var_dir / f'woa23_{var_short}{month}.nc'

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

    log.info(f'[DOWNLOAD] {label}')
    url = build_url(var_long, var_short, month, grid, grid_code)
    return fetch_with_progress(url, dest)


def verify_mode(grid, grid_code):
    """Download only the annual T file to confirm connectivity."""
    log.info(f'=== VERIFY: WOA23 annual temperature, {grid}° resolution ===')
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    ok = download_one('temperature', 't', '00', OUTPUT_DIR, grid, grid_code)
    if ok:
        log.info('=== VERIFICATION PASSED ===')
    else:
        log.error('=== VERIFICATION FAILED ===')
        sys.exit(1)


def main():
    parser = argparse.ArgumentParser(description='Download WOA23 climatology')
    parser.add_argument('--verify', action='store_true',
                        help='Download annual T file only to test connectivity')
    parser.add_argument('--hires', action='store_true',
                        help='Use 0.25° resolution (~1 GB/file) instead of 1°')
    args = parser.parse_args()

    grid, grid_code = GRID_025DEG if args.hires else GRID_1DEG
    res_label = '0.25°' if args.hires else '1.00°'
    log.info(f'Resolution: {res_label}  (use --hires for 0.25°)')

    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)

    fh = logging.FileHandler(OUTPUT_DIR / 'download_woa23.log')
    fh.setFormatter(logging.Formatter('%(asctime)s [%(levelname)s] %(message)s'))
    log.addHandler(fh)

    if args.verify:
        verify_mode(grid, grid_code)
        return

    failed = []
    for var_long, var_short in WOA_VARS.items():
        var_dir = OUTPUT_DIR / var_long
        var_dir.mkdir(exist_ok=True)
        for month in MONTHS:
            if not download_one(var_long, var_short, month,
                                var_dir, grid, grid_code):
                failed.append(f'{var_long}/{month}')

    n_total = len(WOA_VARS) * len(MONTHS)
    if failed:
        log.warning(f'Failed ({len(failed)}/{n_total}): {failed}')
    else:
        log.info(f'All {n_total} WOA23 files downloaded successfully.')


if __name__ == '__main__':
    main()
