#!/usr/bin/env python3
import argparse
import csv
import io
import re
import subprocess
from pathlib import Path
from typing import Iterable, List, Optional, Tuple
from urllib.parse import unquote, urlparse


DEFAULT_CSV = Path(r"/hdd1b/cryosparc/nbdbStatic/script/emdb_validation.csv")
DEFAULT_BASE_DIR = Path("/hdd1b/cryosparc/nbdbStatic/EMDB")
DEFAULT_FAIL_LOG = Path(
    "/hdd1b/cryosparc/nbdbStatic/script/emdb-validation-download-fail.csv"
)
MAX_DATA_ROWS = 60500
# 英文逗号和中文逗号都视为链接分隔符，避免连续链接被识别成一个地址。
HTTPS_PATTERN = re.compile(r"https://[^\s<>\"',，]+", re.IGNORECASE)


def read_csv_text(path: Path) -> Tuple[str, str]:
    """使用常见编码读取 CSV，避免非 UTF-8 文件直接报错。"""
    raw_data = path.read_bytes()
    for encoding in ("utf-8-sig", "utf-8", "gb18030", "cp1252"):
        try:
            return raw_data.decode(encoding), encoding
        except UnicodeDecodeError:
            continue
    raise UnicodeError("无法识别 CSV 文件编码: {}".format(path))


def create_reader(csv_text: str):
    """自动识别逗号、分号或制表符分隔格式。"""
    sample = csv_text[:8192]
    try:
        dialect = csv.Sniffer().sniff(sample, delimiters=",;\t")
    except csv.Error:
        dialect = csv.excel
    return csv.reader(io.StringIO(csv_text), dialect)


def extract_emdb_id(value: str) -> Optional[str]:
    """从第一列提取连续数字，例如 EMD-45218 转换为 45218。"""
    match = re.search(r"\d+", str(value).strip())
    return match.group(0) if match else None


def extract_https_links(cells: Iterable[str]) -> List[str]:
    """从多个单元格中提取全部 HTTPS 链接，并按出现顺序去重。"""
    links = []
    seen = set()
    for cell in cells:
        for match in HTTPS_PATTERN.findall(str(cell)):
            link = match.rstrip(".,;:!?)]}，；。！？」】")
            if link and link not in seen:
                seen.add(link)
                links.append(link)
    return links


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="读取 CSV 前 10 条数据记录并下载 EMDB 验证文件"
    )
    parser.add_argument("--csv", type=Path, default=DEFAULT_CSV, help="输入 CSV 文件")
    parser.add_argument(
        "--base-dir", type=Path, default=DEFAULT_BASE_DIR, help="下载根目录"
    )
    parser.add_argument(
        "--fail-log", type=Path, default=DEFAULT_FAIL_LOG, help="失败日志路径"
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    csv_text, encoding = read_csv_text(args.csv)
    reader = create_reader(csv_text)
    next(reader, None)  # 跳过第一行表头

    failures = []
    processed_links = 0

    for data_row_number, row in enumerate(reader, start=1):
        if data_row_number > MAX_DATA_ROWS:
            break
        if not row:
            print("第 {} 条记录为空，跳过".format(data_row_number))
            continue

        emdb_id = extract_emdb_id(row[0])
        if emdb_id is None:
            print("第 {} 条记录无法提取数字 ID，跳过: {}".format(data_row_number, row[0]))
            continue

        links = extract_https_links(row[1:8])
        if not links:
            print("EMDB {} 的第 2-8 列没有 HTTPS 链接，跳过".format(emdb_id))
            continue

        target_dir = args.base_dir / emdb_id
        target_dir.mkdir(parents=True, exist_ok=True)

        for link in links:
            file_name = unquote(Path(urlparse(link).path).name)
            target_file = target_dir / file_name if file_name else None
            if target_file is not None and target_file.is_file():
                print("文件已存在，跳过: {}".format(target_file))
                continue

            command = ["wget", "-P", str(target_dir), link]
            print("执行: {}".format(" ".join(command)))
            result = subprocess.run(command, check=False)
            processed_links += 1
            if result.returncode != 0:
                failures.append((emdb_id, link))
                print(
                    "下载失败，继续下一条: EMDB {}, 返回码 {}".format(
                        emdb_id, result.returncode
                    )
                )

    args.fail_log.parent.mkdir(parents=True, exist_ok=True)
    with args.fail_log.open("w", encoding="utf-8-sig", newline="") as log_file:
        writer = csv.writer(log_file)
        writer.writerow(["emdb-id", "下载失败的提取链接"])
        writer.writerows(failures)

    print("CSV 编码: {}".format(encoding))
    print("已检查前 {} 条数据记录".format(MAX_DATA_ROWS))
    print("共执行 {} 个下载，失败 {} 个".format(processed_links, len(failures)))
    print("失败日志: {}".format(args.fail_log))


if __name__ == "__main__":
    main()
