# -*- coding: utf-8 -*-
"""
HYCOM 再分析数据下载脚本（1997-2024）
via OPeNDAP + netCDF4（按变量/区域裁剪）

变量:   water_temp, salinity, water_u, water_v, surf_el
区域:   100E-126E, 10N-32N
时间:   1997-01-01 ~ 2024-12-31 (日均)
输出:   每天一个 NetCDF 文件

━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
数据集分段说明（按时间顺序，本脚本采用如下优先规则）：

【再分析段 Reanalysis, GLBv0.08 = 1/12°】
  expt_53.X : 1994-01-01 ~ 2015-12-30
              URL: .../GLBv0.08/expt_53.X/data/{year}

  expt_56.3 : 1994-01-01 ~ 2015-12-31   ← 56.3 与 53.X 重叠，本脚本不用
  expt_57.2 : 2014-07-01 ~ 2016-09-30
              URL: .../GLBv0.08/expt_57.2/data/{year}
  expt_92.8 : 2016-05-01 ~ 2017-01-31
              URL: .../GLBv0.08/expt_92.8/data/{year}
  expt_57.7 : 2017-02-01 ~ 2017-05-31
              URL: .../GLBv0.08/expt_57.7/data/{year}
  expt_92.9 : 2017-06-01 ~ 2017-09-30   ← 极短，可用 93.0 替代
              URL: .../GLBv0.08/expt_92.9/data/{year}
  expt_93.0 : 2018-01-01 ~ 2020-02-18
              URL: .../GLBv0.08/expt_93.0/data/{year}

【分析段 Analysis, GLBy0.08 = 1/25°，分辨率提升】
  expt_93.0 : 2020-02-19 ~ 至今（约到 2024 年底）
              URL: .../GLBy0.08/expt_93.0/data/{year}

本脚本采用的分段映射（唯一、无重叠）：
  1997-01-01 ~ 2015-12-30 → GLBv0.08/expt_53.X
  2016-01-01 ~ 2016-04-30 → GLBv0.08/expt_57.2  (57.2 覆盖至 2016-09-30)
  2016-05-01 ~ 2017-01-31 → GLBv0.08/expt_92.8
  2017-02-01 ~ 2017-05-31 → GLBv0.08/expt_57.7
  2017-06-01 ~ 2017-12-31 → GLBv0.08/expt_92.9
  2018-01-01 ~ 2020-02-18 → GLBv0.08/expt_93.0
  2020-02-19 ~ 2024-12-31 → GLBy0.08/expt_93.0  (1/25° 高分辨率)

注：2015-12-31 在 expt_53.X 中可能缺失，已自动 fallback 到 expt_57.2。
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
"""

import os
import sys
import time
import logging
from datetime import date, timedelta

import numpy as np
import netCDF4 as nc
import cftime

# ─────────────────────────────────────────────
# 用户配置区
# ─────────────────────────────────────────────
LON_MIN, LON_MAX = 100.0, 126.0
LAT_MIN, LAT_MAX = 10.0,  32.0
DATE_START = date(2013, 1, 1)
DATE_END   = date(2024, 12, 31)

OUTPUT_DIR = "./hycom_data"
RETRY_MAX  = 3      # 整体（单日）重试次数
RETRY_WAIT = 15     # 重试等待秒数

# 单次 OPeNDAP 请求按深度分块，避免数据包过大被截断
# 40 层全深度时建议 5~10；若仍报 "packet too short" 可继续调小
DEPTH_CHUNK = 5
CHUNK_RETRY = 4     # 单个数据块的重试次数
CHUNK_WAIT  = 8     # 数据块重试等待秒数

VARIABLES = ["water_temp", "salinity", "water_u", "water_v", "surf_el"]

THREDDS_ROOT = "https://tds.hycom.org/thredds/dodsC"

# ─────────────────────────────────────────────
# 数据集分段表（按时间顺序，无重叠）
# 每条: (start, end, url_template)
# url_template 中 {year} 会被替换为目标年份
# ─────────────────────────────────────────────
# 注意：GLBv0.08（1/12°）与 GLBy0.08（1/25°）网格不同，
#       2020-02-19 后切换到更高分辨率的 GLBy0.08。
DATASET_SEGMENTS = [
    # (date_start, date_end, url_template, grid_label)
    (
        date(1994,  1,  1), date(2015, 12, 30),
        f"{THREDDS_ROOT}/GLBv0.08/expt_53.X/data/{{year}}",
        "GLBv0.08/expt_53.X"
    ),
    (
        date(2016,  1,  1), date(2016,  4, 30),
        f"{THREDDS_ROOT}/GLBv0.08/expt_57.2/data/{{year}}",
        "GLBv0.08/expt_57.2"
    ),
    (
        date(2016,  5,  1), date(2017,  1, 31),
        f"{THREDDS_ROOT}/GLBv0.08/expt_92.8/data/{{year}}",
        "GLBv0.08/expt_92.8"
    ),
    (
        date(2017,  2,  1), date(2017,  5, 31),
        f"{THREDDS_ROOT}/GLBv0.08/expt_57.7/data/{{year}}",
        "GLBv0.08/expt_57.7"
    ),
    (
        date(2017,  6,  1), date(2017, 12, 31),
        f"{THREDDS_ROOT}/GLBv0.08/expt_92.9/data/{{year}}",
        "GLBv0.08/expt_92.9"
    ),
    (
        date(2018,  1,  1), date(2020,  2, 18),
        f"{THREDDS_ROOT}/GLBv0.08/expt_93.0/data/{{year}}",
        "GLBv0.08/expt_93.0"
    ),
    (
        date(2020,  2, 19), date(2024, 12, 31),
        f"{THREDDS_ROOT}/GLBy0.08/expt_93.0/data/{{year}}",
        "GLBy0.08/expt_93.0"   # 1/25° 高分辨率
    ),
]

# ─────────────────────────────────────────────
# 日志设置（Windows GBK 兼容，强制 UTF-8）
# ─────────────────────────────────────────────
os.makedirs(OUTPUT_DIR, exist_ok=True)

log_formatter = logging.Formatter("%(asctime)s [%(levelname)s] %(message)s")

console_handler = logging.StreamHandler(
    stream=open(sys.stdout.fileno(), mode="w", encoding="utf-8", closefd=False)
)
console_handler.setFormatter(log_formatter)

file_handler = logging.FileHandler(
    os.path.join(OUTPUT_DIR, "download.log"), encoding="utf-8"
)
file_handler.setFormatter(log_formatter)

logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)
logger.addHandler(console_handler)
logger.addHandler(file_handler)


# ─────────────────────────────────────────────
# 数据集路由
# ─────────────────────────────────────────────
def get_dataset_info(target_date: date):
    """
    根据目标日期返回 (url, grid_label)。
    遍历 DATASET_SEGMENTS 找到匹配段；找不到则抛出 ValueError。
    """
    for seg_start, seg_end, url_tpl, label in DATASET_SEGMENTS:
        if seg_start <= target_date <= seg_end:
            url = url_tpl.format(year=target_date.year)
            return url, label
    raise ValueError(
        f"日期 {target_date} 未落入任何已配置数据集段，"
        f"当前覆盖范围: {DATASET_SEGMENTS[0][0]} ~ {DATASET_SEGMENTS[-1][1]}"
    )


# ─────────────────────────────────────────────
# 工具函数
# ─────────────────────────────────────────────
def find_range_indices(values, vmin, vmax):
    """返回覆盖 [vmin, vmax] 的索引范围（左闭右开切片）"""
    idx = np.where((values >= vmin) & (values <= vmax))[0]
    if len(idx) == 0:
        raise ValueError(f"坐标范围 [{vmin}, {vmax}] 在数据中无匹配网格点")
    return int(idx[0]), int(idx[-1]) + 1


def get_time_index(remote_ds, target_date: date) -> int:
    """在已打开的 netCDF4 数据集中查找目标日期的时间索引"""
    tv = remote_ds.variables["time"]
    cal = getattr(tv, "calendar", "standard")
    times = cftime.num2date(tv[:], tv.units, calendar=cal)
    for i, t in enumerate(times):
        if (t.year == target_date.year and
                t.month == target_date.month and
                t.day == target_date.day):
            return i
    raise ValueError(
        f"数据集中找不到日期 {target_date}，"
        f"可用范围: {times[0]} ~ {times[-1]}"
    )


def _read_block(src, slices, label):
    """读取一个数据块，带独立重试。slices 为索引元组。"""
    last_err = None
    for k in range(1, CHUNK_RETRY + 1):
        try:
            return src[slices]
        except Exception as e:
            last_err = e
            logger.warning(f"    [{label}] 块读取第 {k} 次失败: {e}")
            if k < CHUNK_RETRY:
                time.sleep(CHUNK_WAIT)
    raise RuntimeError(f"块 [{label}] 读取彻底失败: {last_err}")


def read_3d_var(src, t_idx, lat0, lat1, lon0, lon1, n_depth, vname):
    """按深度分块读取 (time, depth, lat, lon) 变量，拼接成完整数组"""
    parts = []
    for d0 in range(0, n_depth, DEPTH_CHUNK):
        d1 = min(d0 + DEPTH_CHUNK, n_depth)
        label = f"{vname} depth[{d0}:{d1}]"
        block = _read_block(
            src, (t_idx, slice(d0, d1), slice(lat0, lat1), slice(lon0, lon1)),
            label
        )
        parts.append(np.asarray(block))
    return np.concatenate(parts, axis=0)  # 沿 depth 轴拼接


def read_2d_var(src, t_idx, lat0, lat1, lon0, lon1, vname):
    """读取 (time, lat, lon) 变量，如 surf_el"""
    label = f"{vname} surface"
    block = _read_block(
        src, (t_idx, slice(lat0, lat1), slice(lon0, lon1)), label
    )
    return np.asarray(block)


# ─────────────────────────────────────────────
# 核心下载函数
# ─────────────────────────────────────────────
def download_one_day(target_date: date) -> bool:
    """下载单日数据并保存为 NetCDF4 文件，支持断点续传和重试"""
    out_fname = os.path.join(
        OUTPUT_DIR, f"hycom_{target_date.strftime('%Y%m%d')}.nc"
    )
    if os.path.exists(out_fname):
        logger.info(f"已存在，跳过: {out_fname}")
        return True

    url, grid_label = get_dataset_info(target_date)

    for attempt in range(1, RETRY_MAX + 1):
        try:
            logger.info(
                f"[{target_date}] 数据集={grid_label}  "
                f"URL={url}  (第 {attempt} 次尝试)"
            )

            remote = nc.Dataset(url)

            lon_vals   = remote.variables["lon"][:]
            lat_vals   = remote.variables["lat"][:]
            depth_vals = remote.variables["depth"][:]

            lon_i0, lon_i1 = find_range_indices(lon_vals, LON_MIN, LON_MAX)
            lat_i0, lat_i1 = find_range_indices(lat_vals, LAT_MIN, LAT_MAX)
            t_idx = get_time_index(remote, target_date)

            lon_sub = lon_vals[lon_i0:lon_i1]
            lat_sub = lat_vals[lat_i0:lat_i1]

            logger.info(
                f"  t_idx={t_idx}, "
                f"lon [{lon_sub[0]:.2f}~{lon_sub[-1]:.2f}] ({len(lon_sub)} pts) "
                f"lat [{lat_sub[0]:.2f}~{lat_sub[-1]:.2f}] ({len(lat_sub)} pts) "
                f"depth {len(depth_vals)} layers"
            )

            with nc.Dataset(out_fname, "w", format="NETCDF4") as ds:
                ds.title       = f"HYCOM {grid_label} Reanalysis/Analysis"
                ds.date        = str(target_date)
                ds.dataset     = grid_label
                ds.lon_range   = f"{LON_MIN}E-{LON_MAX}E"
                ds.lat_range   = f"{LAT_MIN}N-{LAT_MAX}N"
                ds.source_url  = url
                ds.Conventions = "CF-1.6"

                ds.createDimension("time",  1)
                ds.createDimension("depth", len(depth_vals))
                ds.createDimension("lat",   len(lat_sub))
                ds.createDimension("lon",   len(lon_sub))

                vt          = ds.createVariable("time", "f8", ("time",))
                vt.units    = "days since 1900-01-01 00:00:00"
                vt.calendar = "gregorian"
                vt[:]       = cftime.date2num(
                    cftime.datetime(target_date.year,
                                    target_date.month,
                                    target_date.day),
                    vt.units
                )

                vd          = ds.createVariable("depth", "f4", ("depth",))
                vd.units    = "m"
                vd.positive = "down"
                vd[:]       = depth_vals

                vlat           = ds.createVariable("lat", "f4", ("lat",))
                vlat.units     = "degrees_north"
                vlat.long_name = "latitude"
                vlat[:]        = lat_sub

                vlon           = ds.createVariable("lon", "f4", ("lon",))
                vlon.units     = "degrees_east"
                vlon.long_name = "longitude"
                vlon[:]        = lon_sub

                for vname in VARIABLES:
                    if vname not in remote.variables:
                        logger.warning(f"  变量 {vname} 不在数据集中，跳过")
                        continue

                    src  = remote.variables[vname]
                    fill = getattr(src, "_FillValue", 1.0e30)

                    if vname == "surf_el":
                        logger.info(f"  读取 {vname} ...")
                        data = read_2d_var(
                            src, t_idx, lat_i0, lat_i1, lon_i0, lon_i1, vname
                        )
                        v = ds.createVariable(
                            vname, "f4", ("time", "lat", "lon"),
                            fill_value=fill, zlib=True, complevel=4
                        )
                        v[:] = data[np.newaxis, ...]
                    else:
                        logger.info(
                            f"  读取 {vname} (按 {DEPTH_CHUNK} 层分块) ..."
                        )
                        data = read_3d_var(
                            src, t_idx, lat_i0, lat_i1, lon_i0, lon_i1,
                            len(depth_vals), vname
                        )
                        v = ds.createVariable(
                            vname, "f4", ("time", "depth", "lat", "lon"),
                            fill_value=fill, zlib=True, complevel=4
                        )
                        v[:] = data[np.newaxis, ...]

                    for attr in ("long_name", "units", "standard_name",
                                 "scale_factor", "add_offset",
                                 "valid_min", "valid_max"):
                        if hasattr(src, attr):
                            setattr(v, attr, getattr(src, attr))

            remote.close()
            logger.info(f"  [OK] 已保存: {out_fname}")
            return True

        except Exception as e:
            logger.error(f"  [FAIL] 第 {attempt} 次失败: {e}")
            if os.path.exists(out_fname):
                os.remove(out_fname)
            if attempt < RETRY_MAX:
                logger.info(f"  等待 {RETRY_WAIT}s 后重试...")
                time.sleep(RETRY_WAIT)
            else:
                logger.error(f"  放弃 {target_date}，已记录到失败列表")
                fail_log = os.path.join(OUTPUT_DIR, "failed_dates.txt")
                with open(fail_log, "a", encoding="utf-8") as f:
                    f.write(f"{target_date}\t{grid_label}\t{e}\n")
                return False

    return False


# ─────────────────────────────────────────────
# 主流程
# ─────────────────────────────────────────────
def main():
    total_days = (DATE_END - DATE_START).days + 1
    logger.info("=" * 70)
    logger.info("HYCOM 数据下载任务启动（1997-2024 多数据集版）")
    logger.info(f"共需下载 {total_days} 天: {DATE_START} ~ {DATE_END}")
    logger.info(f"空间范围: {LON_MIN}E-{LON_MAX}E, {LAT_MIN}N-{LAT_MAX}N")
    logger.info(f"变量列表: {VARIABLES}")
    logger.info(f"输出目录: {os.path.abspath(OUTPUT_DIR)}")
    logger.info("数据集分段（唯一映射）:")
    for seg_start, seg_end, _, label in DATASET_SEGMENTS:
        logger.info(f"  {seg_start} ~ {seg_end}  →  {label}")
    logger.info("=" * 70)

    success, failed = 0, 0
    cur = DATE_START
    while cur <= DATE_END:
        ok = download_one_day(cur)
        if ok:
            success += 1
        else:
            failed += 1
        cur += timedelta(days=1)
        time.sleep(0.5)

    logger.info("=" * 70)
    logger.info(f"下载完成！成功: {success} 天，失败: {failed} 天")
    if failed:
        logger.info(
            f"失败日期见: {os.path.join(OUTPUT_DIR, 'failed_dates.txt')}"
        )
    logger.info("=" * 70)


if __name__ == "__main__":
    main()
