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': 'u10',   'desc': '10m U wind'},
    'v10m':  {'code': '128_166', 'fname': 'v10',   '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'},
    'rsn':   {'code': '128_033', 'fname': 'rsn',    'desc': 'Sea-ice cover'},
}
BASE_URL = "https://nsf-ncar-era5.s3.amazonaws.com"

# =============================================================================
# 2. URL 构建函数（返回url + 云端原始文件名）
# =============================================================================
def build_pl_url(var_key: str, date: datetime) -> tuple[str, str]:
    cfg = PL_VARS[var_key]
    yyyymm = date.strftime('%Y%m')
    yyyymmdd = date.strftime('%Y%m%d')
    # u/v 使用 ll025uv，其余变量 ll025sc
    suffix = "ll025uv" if var_key in ("u", "v") else "ll025sc"
    # 云端原始文件名
    file_name = (
        f"e5.oper.an.pl.{cfg['code']}_{cfg['fname']}.{suffix}."
        f"{yyyymmdd}00_{yyyymmdd}23.nc"
    )
    url = f"{BASE_URL}/e5.oper.an.pl/{yyyymm}/{file_name}"
    return url, file_name

def build_sfc_url(var_key: str, year: int, month: int) -> tuple[str, str]:
    cfg = SFC_VARS[var_key]
    yyyymm = f"{year}{month:02d}"
    last_day = calendar.monthrange(year, month)[1]
    file_name = (
        f"e5.oper.an.sfc.{cfg['code']}_{cfg['fname']}.ll025sc."
        f"{yyyymm}0100_{yyyymm}{last_day}23.nc"
    )
    url = f"{BASE_URL}/e5.oper.an.sfc/{yyyymm}/{file_name}"
    return url, file_name

# =============================================================================
# 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:
    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:
                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]]:
    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, fname = build_pl_url(v, current)
            yyyymm = current.strftime('%Y%m')
            # 按月文件夹存放，不按天分文件夹
            month_dir = os.path.join(out_dir, yyyymm)
            out_path = os.path.join(month_dir, fname)
            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]]:
    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, fname = build_sfc_url(v, current.year, current.month)
            yyyymm = current.strftime('%Y%m')
            month_dir = os.path.join(out_dir, yyyymm)
            out_path = os.path.join(month_dir, fname)
            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]]:
    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=['z', 't', 'u', 'v', 'q','r'],
    #     start_date='2024-12-01',
    #     end_date='2024-12-31',
    #     out_dir='era5_aws_pl/',
    #     max_workers=4
    #)
    # 生成目录结构示例：
    # era5_aws_pl/202401/e5.oper.an.pl.128_131_u.ll025uv.2024010100_2024010123.nc
    # era5_aws_pl/202401/e5.oper.an.pl.128_129_z.ll025sc.2024010100_2024010123.nc

    # ---------------- 示例2：地表层逐月 ----------------
    download_sfc_range(
         var_list=['rsn','t2m', 'd2m', 'u10m', 'v10m', 'msl','sp','swvl1','swvl2','swvl3','swvl4','stl1','stl2','stl3','stl4','sst','sk','sd','ci'],
         start_year=2024, start_month=1,
         end_year=2024, end_month=12,
         out_dir='era5_aws_sfc/',
         max_workers=4
    )

    # ---------------- 示例3：混合变量统一入口 ----------------
    # download_era5(
    #    var_list=['z', 't', 'u', 'v', 'msl', 't2m'],
    #    start='2024-01-01',
    #    end='2024-01-31',
    #    out_dir='era5_aws/',
    #    max_workers=4,
    #    retries=5
    # )

