import os
import sys
import time
import calendar
from datetime import datetime, timedelta
from concurrent.futures import ThreadPoolExecutor, as_completed
from urllib.request import urlopen, Request
from urllib.error import HTTPError, URLError


# =============================================================================
# 1. 变量配置表
# =============================================================================

# 气压层变量（逐日文件）
PL_VARS = {
    'z':  {'code': '128_129', 'fname': 'z',   'desc': 'Geopotential'},
    't':  {'code': '128_130', 'fname': 't',   'desc': 'Temperature'},
    'u':  {'code': '128_131', 'fname': 'u',   'desc': 'U wind'},
    'v':  {'code': '128_132', 'fname': 'v',   'desc': 'V wind'},
    'q':  {'code': '128_133', 'fname': 'q',   'desc': 'Specific humidity'},
    'r':  {'code': '128_157', 'fname': 'r',   'desc': 'Relative_humidity'},
}

# 地表层变量（逐月文件）
SFC_VARS = {
    'msl':   {'code': '128_151', 'fname': 'msl',   'desc': 'Mean sea level pressure'},
    't2m':   {'code': '128_167', 'fname': '2t',    'desc': '2m temperature'},
    'd2m':   {'code': '128_168', 'fname': '2d',    'desc': '2m dewpoint temperature'},
    'sp':  {'code': '128_134', 'fname': 'sp',  'desc': 'surface_pressure'},
    'u10m':  {'code': '128_165', 'fname': '10u',  'desc': '10m U wind'},
    'v10m':  {'code': '128_166', 'fname': '10v',  'desc': '10m V wind'},
    'swvl1':  {'code': '128_039', 'fname': 'swvl1',  'desc': 'volumetric_soil_water_layer_1'},
    'swvl2':  {'code': '128_040', 'fname': 'swvl2',  'desc': 'volumetric_soil_water_layer_2'},
    'swvl3':  {'code': '128_041', 'fname': 'swvl3',  'desc': 'volumetric_soil_water_layer_3'},
    'swvl4':  {'code': '128_042', 'fname': 'swvl4',  'desc': 'volumetric_soil_water_layer_4'},
    'stl1':  {'code': '128_139', 'fname': 'stl1',  'desc': 'soil_temperature_level_1'},
    'stl2':  {'code': '128_170', 'fname': 'stl2',  'desc': 'soil_temperature_level_2'},
    'stl3':  {'code': '128_183', 'fname': 'stl3',  'desc': 'soil_temperature_level_3'},
    'stl4':  {'code': '128_236', 'fname': 'stl4',  'desc': 'soil_temperature_level_4'},
    'sst':  {'code': '128_034', 'fname': 'sstk',  'desc': 'Sea surface temperature'},
    'sk':  {'code': '128_235', 'fname': 'skt',  'desc': 'Skin temperature'},
    'sd':  {'code': '128_141', 'fname': 'sd',  'desc': 'Snow depth'},
    'ci':  {'code': '128_031', 'fname': 'ci',  'desc': 'Sea-ice cover'},
}

BASE_URL = "https://nsf-ncar-era5.s3.amazonaws.com"

#BASE_URL = "https://nsf-ncar-era5.s3.amazonaws.com/e5.oper.an.sfc/"

# =============================================================================
# 2. URL 构建函数
# =============================================================================

def build_pl_url(var_key: str, date: datetime) -> str:
    """
    构建气压层（逐日）变量的下载 URL

    Parameters
    ----------
    var_key : str
        变量名，如 'z', 't', 'u', 'v', 'q'
    date : datetime
        日期（年月日）

    Returns
    -------
    str
        完整的下载 URL
    """
    cfg = PL_VARS[var_key]
    yyyymm = date.strftime('%Y%m')
    yyyymmdd = date.strftime('%Y%m%d')

    url = (
        f"{BASE_URL}/e5.oper.an.pl/{yyyymm}/"
        f"e5.oper.an.pl.{cfg['code']}_{cfg['fname']}.ll025uv."
        f"{yyyymmdd}00_{yyyymmdd}23.nc"
    )
    return url


def build_sfc_url(var_key: str, year: int, month: int) -> str:
    """
    构建地表层（逐月）变量的下载 URL

    Parameters
    ----------
    var_key : str
        变量名，如 'msl', 't2m', 'u10m' 等
    year : int
        年份
    month : int
        月份

    Returns
    -------
    str
        完整的下载 URL
    """
    cfg = SFC_VARS[var_key]
    yyyymm = f"{year}{month:02d}"
    last_day = calendar.monthrange(year, month)[1]

    url = (
        f"{BASE_URL}/e5.oper.an.sfc/{yyyymm}/"
        f"e5.oper.an.sfc.{cfg['code']}_{cfg['fname']}.ll025sc."
        f"{yyyymm}0100_{yyyymm}{last_day}23.nc"
    )
    return url


# =============================================================================
# 3. 核心下载函数（支持断点续传 + 重试）
# =============================================================================

def download_file(url: str, out_path: str,
                  retries: int = 5,
                  backoff_factor: float = 2.0,
                  timeout: int = 120,
                  chunk_size: int = 8192,
                  verbose: bool = True) -> bool:
    """
    下载单个文件，支持断点续传和指数退避重试

    Parameters
    ----------
    url : str
        下载地址
    out_path : str
        本地保存路径
    retries : int
        最大重试次数
    backoff_factor : float
        退避倍数（秒）
    timeout : int
        单次请求超时（秒）
    chunk_size : int
        流式读取块大小
    verbose : bool
        是否打印进度

    Returns
    -------
    bool
        下载是否成功
    """
    os.makedirs(os.path.dirname(out_path), exist_ok=True)

    # 已下载字节数（断点续传）
    existing_size = os.path.getsize(out_path) if os.path.exists(out_path) else 0

    headers = {}
    if existing_size > 0:
        headers['Range'] = f'bytes={existing_size}-'

    for attempt in range(retries + 1):
        try:
            req = Request(url, headers=headers)
            with urlopen(req, timeout=timeout) as response:
                total_size = existing_size + int(response.headers.get('Content-Length', 0))
                mode = 'ab' if existing_size > 0 else 'wb'

                with open(out_path, mode) as f:
                    downloaded = existing_size
                    while True:
                        chunk = response.read(chunk_size)
                        if not chunk:
                            break
                        f.write(chunk)
                        downloaded += len(chunk)

                        if verbose and total_size > 0:
                            pct = downloaded / total_size * 100
                            mb = downloaded / 1024 / 1024
                            total_mb = total_size / 1024 / 1024
                            print(f"\r  ↳ {os.path.basename(out_path):40s} "
                                  f"{pct:5.1f}%  {mb:7.1f}/{total_mb:7.1f} MB", end='', flush=True)

                if verbose:
                    print()  # 换行
                return True

        except HTTPError as e:
            if e.code == 416:  # Range not satisfiable（文件已完整）
                if verbose:
                    print(f"  ↳ {os.path.basename(out_path):40s} already complete.")
                return True
            if e.code == 404:
                print(f"\n[ERROR] 404 Not Found: {url}")
                return False
            print(f"\n[HTTP {e.code}] {url}, attempt {attempt+1}/{retries+1}")

        except (URLError, TimeoutError, ConnectionResetError) as e:
            print(f"\n[Network] {e}, attempt {attempt+1}/{retries+1}")

        if attempt < retries:
            sleep_time = backoff_factor ** attempt
            if verbose:
                print(f"  → Retrying in {sleep_time:.1f}s...")
            time.sleep(sleep_time)
            # 重新获取已下载大小
            existing_size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
            headers['Range'] = f'bytes={existing_size}-'

    print(f"\n[FAILED] Max retries exceeded: {url}")
    return False


# =============================================================================
# 4. 批量下载接口
# =============================================================================

def download_pl_range(var_list: list[str],
                      start_date: str,
                      end_date: str,
                      out_dir: str = "./era5_pl",
                      max_workers: int = 4,
                      retries: int = 5) -> list[tuple[str, bool]]:
    """
    批量下载气压层变量（逐日）

    Parameters
    ----------
    var_list : list[str]
        变量名列表，如 ['z', 't', 'u']
    start_date : str
        起始日期，格式 'YYYY-MM-DD'
    end_date : str
        结束日期，格式 'YYYY-MM-DD'
    out_dir : str
        输出根目录
    max_workers : int
        并发下载线程数
    retries : int
        单个文件重试次数

    Returns
    -------
    list[tuple[str, bool]]
        [(url, success), ...]
    """
    start = datetime.strptime(start_date, '%Y-%m-%d')
    end = datetime.strptime(end_date, '%Y-%m-%d')

    tasks = []
    current = start
    while current <= end:
        for v in var_list:
            if v not in PL_VARS:
                print(f"[WARN] Unknown PL variable: {v}, skipped.")
                continue
            url = build_pl_url(v, current)
            yyyymm = current.strftime('%Y%m')
            yyyymmdd = current.strftime('%Y%m%d')
            out_path = os.path.join(out_dir, v, yyyymm, f"{v}_{yyyymmdd}.nc")
            tasks.append((url, out_path))
        current += timedelta(days=1)

    print(f"[INFO] Total PL files to download: {len(tasks)}")
    results = []

    with ThreadPoolExecutor(max_workers=max_workers) as executor:
        future_to_url = {
            executor.submit(download_file, url, path, retries): url
            for url, path in tasks
        }
        for future in as_completed(future_to_url):
            url = future_to_url[future]
            try:
                ok = future.result()
                results.append((url, ok))
            except Exception as e:
                print(f"\n[EXCEPTION] {url}: {e}")
                results.append((url, False))

    success = sum(1 for _, ok in results if ok)
    print(f"[INFO] PL download complete: {success}/{len(tasks)} succeeded.")
    return results


def download_sfc_range(var_list: list[str],
                       start_year: int,
                       start_month: int,
                       end_year: int,
                       end_month: int,
                       out_dir: str = "./era5_sfc",
                       max_workers: int = 4,
                       retries: int = 5) -> list[tuple[str, bool]]:
    """
    批量下载地表层变量（逐月）

    Parameters
    ----------
    var_list : list[str]
        变量名列表，如 ['msl', 't2m', 'u10m']
    start_year, start_month : int
        起始年月
    end_year, end_month : int
        结束年月
    out_dir : str
        输出根目录
    max_workers : int
        并发下载线程数
    retries : int
        单个文件重试次数

    Returns
    -------
    list[tuple[str, bool]]
        [(url, success), ...]
    """
    tasks = []
    current = datetime(start_year, start_month, 1)
    end = datetime(end_year, end_month, 1)

    while current <= end:
        for v in var_list:
            if v not in SFC_VARS:
                print(f"[WARN] Unknown SFC variable: {v}, skipped.")
                continue
            url = build_sfc_url(v, current.year, current.month)
            yyyymm = current.strftime('%Y%m')
            out_path = os.path.join(out_dir, v, f"{v}_{yyyymm}.nc")
            tasks.append((url, out_path))
        # 下一个月
        if current.month == 12:
            current = datetime(current.year + 1, 1, 1)
        else:
            current = datetime(current.year, current.month + 1, 1)

    print(f"[INFO] Total SFC files to download: {len(tasks)}")
    results = []

    with ThreadPoolExecutor(max_workers=max_workers) as executor:
        future_to_url = {
            executor.submit(download_file, url, path, retries): url
            for url, path in tasks
        }
        for future in as_completed(future_to_url):
            url = future_to_url[future]
            try:
                ok = future.result()
                results.append((url, ok))
            except Exception as e:
                print(f"\n[EXCEPTION] {url}: {e}")
                results.append((url, False))

    success = sum(1 for _, ok in results if ok)
    print(f"[INFO] SFC download complete: {success}/{len(tasks)} succeeded.")
    return results


# =============================================================================
# 5. 便捷封装：按变量类型自动分发
# =============================================================================

def download_era5(var_list: list[str],
                  start: str,
                  end: str,
                  out_dir: str = "./era5",
                  max_workers: int = 4,
                  retries: int = 5) -> list[tuple[str, bool]]:
    """
    通用 ERA5 下载入口，自动识别变量类型并调用对应下载逻辑

    Parameters
    ----------
    var_list : list[str]
        变量名列表，可混合气压层和地表层变量
    start : str
        起始时间。气压层: 'YYYY-MM-DD'；地表层: 'YYYY-MM'
    end : str
        结束时间。气压层: 'YYYY-MM-DD'；地表层: 'YYYY-MM'
    out_dir : str
        输出根目录
    max_workers : int
        并发下载线程数
    retries : int
        单个文件重试次数

    Returns
    -------
    list[tuple[str, bool]]
        [(url, success), ...]
    """
    pl_vars = [v for v in var_list if v in PL_VARS]
    sfc_vars = [v for v in var_list if v in SFC_VARS]
    unknown = [v for v in var_list if v not in PL_VARS and v not in SFC_VARS]

    if unknown:
        print(f"[WARN] Unknown variables skipped: {unknown}")

    all_results = []

    if pl_vars:
        print("=" * 60)
        print(f"[PL ] Downloading: {pl_vars}")
        print(f"[PL ] Date range : {start} ~ {end}")
        print("=" * 60)
        r = download_pl_range(pl_vars, start, end,
                              out_dir=os.path.join(out_dir, "pl"),
                              max_workers=max_workers, retries=retries)
        all_results.extend(r)

    if sfc_vars:
        print("=" * 60)
        print(f"[SFC] Downloading: {sfc_vars}")
        print(f"[SFC] Date range : {start} ~ {end}")
        print("=" * 60)
        # 解析年月
        s_parts = start.split('-')
        e_parts = end.split('-')
        r = download_sfc_range(sfc_vars,
                                 int(s_parts[0]), int(s_parts[1]),
                                 int(e_parts[0]), int(e_parts[1]),
                                 out_dir=os.path.join(out_dir, "sfc"),
                                 max_workers=max_workers, retries=retries)
        all_results.extend(r)

    return all_results


# =============================================================================
# 6. 主程序示例
# =============================================================================

if __name__ == "__main__":
    # ---------------- 示例 1：下载气压层变量（逐日）----------------
    download_pl_range(
         var_list=['u'],
         start_date='2024-01-01',
         end_date='2024-01-31',
         out_dir='era5_aws_pl/',
         max_workers=4
    )

    # ---------------- 示例 2：下载地表层变量（逐月）----------------
    download_sfc_range(
         var_list=['u10m', 'v10m','sp'],
         start_year=2024, start_month=1,
         end_year=2024, end_month=3,
         out_dir='era5_aws_sfc/',
         max_workers=4
    )

    # ---------------- 示例 3：通用接口（混合变量）----------------
    #download_era5(
    #    var_list=['z', 't', 'msl', 't2m', 'u10m'],
    #    start='2024-01-01',   # 气压层按日，地表层会自动取年月
    #    end='2024-01-31',
    #    out_dir='era5_aws/',
    #    max_workers=4,
    #    retries=5
    #)
