#!/usr/bin/env python3
# -*- coding: utf-8 -*-

"""
下载 BWIKI「夜刀神十香 - 碧蓝航线」页面中的所有 MP3，
提取网页显示的中文台词，并使用 faster-whisper 本地识别 MP3 中的日语台词。

输出：
tohka_azurlane/
├─ audio/
│  ├─ 001_自我介绍_xxx.mp3
│  ├─ 002_获取台词_xxx.mp3
│  └─ ...
├─ metadata.csv
└─ metadata.json

metadata.csv 主要字段：
index, category, zh_text, ja_text, filename, url,
asr_model, asr_language, asr_language_probability, asr_status

依赖：
    pip install requests beautifulsoup4 faster-whisper

示例：
    # NVIDIA GPU，推荐
    python download_tohka_azurlane_with_asr.py --device cuda

    # CPU 也可以，只是慢
    python download_tohka_azurlane_with_asr.py --device cpu

    # 如果已经下载好音频，只重新做 ASR
    python download_tohka_azurlane_with_asr.py --device cuda --asr-only

注意：
1. zh_text 是 BWIKI 页面显示的中文翻译。
2. ja_text 是 faster-whisper 根据 MP3 自动识别出的日语原文，可能存在误识别。
3. 真正用于 GPT-SoVITS 训练前，建议人工快速校对 ja_text。
"""

from __future__ import annotations

import argparse
import csv
import json
import re
import time
from pathlib import Path
from urllib.parse import urljoin, urlparse

import requests
from bs4 import BeautifulSoup, NavigableString, Tag
from requests.adapters import HTTPAdapter
from urllib3.util.retry import Retry


PAGE_URL = (
    "https://wiki.biligame.com/blhx/"
    "%E5%A4%9C%E5%88%80%E7%A5%9E%E5%8D%81%E9%A6%99"
)

OUTPUT_DIR = Path("tohka_azurlane")
AUDIO_DIR = OUTPUT_DIR / "audio"
CSV_PATH = OUTPUT_DIR / "metadata.csv"
JSON_PATH = OUTPUT_DIR / "metadata.json"

REQUEST_TIMEOUT = 30
DOWNLOAD_INTERVAL = 0.2

HEADERS = {
    "User-Agent": (
        "Mozilla/5.0 (Windows NT 10.0; Win64; x64) "
        "AppleWebKit/537.36 (KHTML, like Gecko) "
        "Chrome/152.0.0.0 Safari/537.36"
    ),
    "Referer": PAGE_URL,
}

CSV_FIELDS = [
    "index",
    "category",
    "zh_text",
    "ja_text",
    "filename",
    "url",
    "asr_model",
    "asr_language",
    "asr_language_probability",
    "asr_status",
]


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="下载夜刀神十香 BWIKI 语音，并用 faster-whisper 转写日语。"
    )
    parser.add_argument(
        "--model",
        default="large-v3",
        help="Whisper 模型名称，默认：large-v3",
    )
    parser.add_argument(
        "--device",
        choices=("cuda", "cpu"),
        default="cuda",
        help="ASR 运行设备，默认：cuda",
    )
    parser.add_argument(
        "--compute-type",
        default=None,
        help=(
            "CTranslate2 计算类型。默认根据设备自动选择："
            "cuda=int8_float16，cpu=int8"
        ),
    )
    parser.add_argument(
        "--beam-size",
        type=int,
        default=5,
        help="Whisper beam size，默认：5",
    )
    parser.add_argument(
        "--asr-only",
        action="store_true",
        help="不重新抓网页/下载，只对已有 metadata + audio 执行 ASR。",
    )
    parser.add_argument(
        "--force-asr",
        action="store_true",
        help="即使已有 ja_text，也重新识别。",
    )
    return parser.parse_args()


def make_session() -> requests.Session:
    session = requests.Session()

    retry = Retry(
        total=4,
        connect=4,
        read=4,
        status=4,
        backoff_factor=0.8,
        status_forcelist=(429, 500, 502, 503, 504),
        allowed_methods=frozenset({"GET"}),
    )

    session.mount("https://", HTTPAdapter(max_retries=retry))
    session.headers.update(HEADERS)
    return session


def clean_text(text: str) -> str:
    text = text.replace("\xa0", " ")
    text = re.sub(r"[ \t\r\f\v]+", " ", text)
    text = re.sub(r"\s*\n\s*", "\n", text)
    return text.strip(" \n|")


def clean_asr_text(text: str) -> str:
    """
    对 Whisper 输出做轻量清理。
    不擅自替换日文内容，只整理多余空白。
    """
    text = text.replace("\u3000", " ")
    text = re.sub(r"\s+", " ", text).strip()

    # Whisper 有时会在日文标点两侧产生空格。
    text = re.sub(r"\s+([、。！？!?」』）】])", r"\1", text)
    text = re.sub(r"([「『（【])\s+", r"\1", text)

    return text


def safe_filename(text: str, max_len: int = 35) -> str:
    text = clean_text(text)
    text = re.sub(r'[\\/:*?"<>|]', "_", text)
    text = re.sub(r"\s+", "_", text)
    text = text.strip("._ ")
    return (text[:max_len] or "voice")


def get_mp3_url(tag: Tag, base_url: str) -> str | None:
    for attr in ("href", "src", "data-src", "data-url"):
        value = tag.get(attr)
        if not value:
            continue

        value = str(value).strip()
        path = urlparse(value).path.lower()

        if path.endswith(".mp3"):
            return urljoin(base_url, value)

    return None


def walk_tokens(node, base_url: str):
    """
    按 HTML 出现顺序产出：
      ("text", 文本)
      ("audio", mp3_url)
    """
    if isinstance(node, NavigableString):
        yield ("text", str(node))
        return

    if not isinstance(node, Tag):
        return

    if node.name in {"script", "style", "noscript"}:
        return

    mp3_url = get_mp3_url(node, base_url)
    if mp3_url:
        yield ("audio", mp3_url)
        return

    if node.name in {"br", "p", "div", "li"}:
        yield ("text", "\n")

    for child in node.children:
        yield from walk_tokens(child, base_url)

    if node.name in {"p", "div", "li"}:
        yield ("text", "\n")


def row_contains_mp3(row: Tag, base_url: str) -> bool:
    for tag in row.find_all(True):
        if get_mp3_url(tag, base_url):
            return True
    return False


def extract_records(html: str, base_url: str) -> list[dict]:
    soup = BeautifulSoup(html, "html.parser")

    records: list[dict] = []
    seen_urls: set[str] = set()

    for row in soup.find_all("tr"):
        if not row_contains_mp3(row, base_url):
            continue

        cells = row.find_all(["th", "td"], recursive=False)

        if len(cells) >= 2:
            category = clean_text(cells[0].get_text(" ", strip=True))
            content_cells = cells[1:]
        else:
            category = ""
            content_cells = [row]

        for cell in content_cells:
            text_buffer: list[str] = []

            for token_type, value in walk_tokens(cell, base_url):
                if token_type == "text":
                    text_buffer.append(value)
                    continue

                zh_text = clean_text("".join(text_buffer))
                text_buffer.clear()

                if value in seen_urls:
                    continue

                seen_urls.add(value)

                records.append(
                    {
                        "category": category,
                        "zh_text": zh_text,
                        "ja_text": "",
                        "url": value,
                        "asr_model": "",
                        "asr_language": "",
                        "asr_language_probability": "",
                        "asr_status": "pending",
                    }
                )

    return records


def download_file(
    session: requests.Session,
    url: str,
    path: Path,
) -> None:
    if path.exists() and path.stat().st_size > 0:
        print(f"[跳过] 已存在: {path.name}")
        return

    with session.get(url, timeout=REQUEST_TIMEOUT, stream=True) as response:
        response.raise_for_status()

        content_type = response.headers.get("Content-Type", "").lower()
        if content_type and "audio" not in content_type and "mpeg" not in content_type:
            print(
                f"[警告] {path.name} Content-Type={content_type}，仍尝试保存。"
            )

        tmp_path = path.with_suffix(path.suffix + ".part")

        with tmp_path.open("wb") as f:
            for chunk in response.iter_content(chunk_size=1024 * 256):
                if chunk:
                    f.write(chunk)

        tmp_path.replace(path)


def normalize_record(record: dict) -> dict:
    """
    兼容旧版脚本生成的 metadata：
    旧字段 text -> 新字段 zh_text。
    """
    if "zh_text" not in record:
        record["zh_text"] = record.pop("text", "")

    record.setdefault("ja_text", "")
    record.setdefault("asr_model", "")
    record.setdefault("asr_language", "")
    record.setdefault("asr_language_probability", "")
    record.setdefault("asr_status", "pending")

    return record


def save_metadata(records: list[dict]) -> None:
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)

    rows = []
    for raw_record in records:
        record = normalize_record(dict(raw_record))
        rows.append({field: record.get(field, "") for field in CSV_FIELDS})

    # utf-8-sig：Windows Excel 直接打开中文/日文不容易乱码
    with CSV_PATH.open("w", encoding="utf-8-sig", newline="") as f:
        writer = csv.DictWriter(f, fieldnames=CSV_FIELDS)
        writer.writeheader()
        writer.writerows(rows)

    with JSON_PATH.open("w", encoding="utf-8") as f:
        json.dump(rows, f, ensure_ascii=False, indent=2)


def load_existing_metadata() -> list[dict]:
    if JSON_PATH.exists():
        with JSON_PATH.open("r", encoding="utf-8") as f:
            records = json.load(f)
        return [normalize_record(dict(x)) for x in records]

    if CSV_PATH.exists():
        with CSV_PATH.open("r", encoding="utf-8-sig", newline="") as f:
            records = list(csv.DictReader(f))
        return [normalize_record(dict(x)) for x in records]

    raise FileNotFoundError(
        "没有找到 metadata.json 或 metadata.csv。"
        "如果使用 --asr-only，请先运行一次普通下载模式。"
    )


def prepare_records_from_web(session: requests.Session) -> list[dict]:
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    AUDIO_DIR.mkdir(parents=True, exist_ok=True)

    print(f"[1/4] 获取网页: {PAGE_URL}")
    response = session.get(PAGE_URL, timeout=REQUEST_TIMEOUT)
    response.raise_for_status()

    if response.apparent_encoding:
        response.encoding = response.apparent_encoding

    print("[2/4] 解析 MP3 与中文台词...")
    new_records = extract_records(response.text, response.url)

    if not new_records:
        raise RuntimeError(
            "没有解析到任何 MP3。网页结构可能发生变化。"
        )

    # 读取旧 metadata，把已经识别过的日文结果继承回来。
    old_by_url: dict[str, dict] = {}
    if JSON_PATH.exists() or CSV_PATH.exists():
        try:
            old_records = load_existing_metadata()
            old_by_url = {
                r.get("url", ""): r
                for r in old_records
                if r.get("url")
            }
        except Exception as exc:
            print(f"[警告] 旧 metadata 读取失败，将重新生成：{exc}")

    print(f"发现 {len(new_records)} 条不重复 MP3。")

    records: list[dict] = []

    for i, record in enumerate(new_records, start=1):
        url = record["url"]
        category = record["category"] or "未分类"

        original_stem = Path(urlparse(url).path).stem
        filename = (
            f"{i:03d}_"
            f"{safe_filename(category)}_"
            f"{original_stem[:12]}.mp3"
        )

        record["index"] = i
        record["filename"] = filename

        old = old_by_url.get(url)
        if old:
            # 保留旧 ASR 结果，避免重复识别。
            for key in (
                "ja_text",
                "asr_model",
                "asr_language",
                "asr_language_probability",
                "asr_status",
            ):
                if old.get(key) not in (None, ""):
                    record[key] = old.get(key)

        print()
        print(f"[{i:03d}/{len(new_records):03d}] {category}")
        print(f"中文: {record['zh_text'] or '(未提取到文本)'}")
        print(f"URL : {url}")

        try:
            download_file(session, url, AUDIO_DIR / filename)
        except requests.RequestException as exc:
            print(f"[下载失败] {filename}: {exc}")
            record["asr_status"] = "download_failed"

        records.append(record)

        # 每处理一个文件就保存一次，脚本中断也能续跑。
        save_metadata(records)
        time.sleep(DOWNLOAD_INTERVAL)

    print()
    print(f"[完成] 已保存 {CSV_PATH}")
    print(f"[完成] 已保存 {JSON_PATH}")
    return records


def create_whisper_model(
    model_name: str,
    device: str,
    compute_type: str | None,
):
    try:
        from faster_whisper import WhisperModel
    except ImportError as exc:
        raise RuntimeError(
            "未安装 faster-whisper。\n"
            "请执行：pip install faster-whisper"
        ) from exc

    if compute_type is None:
        compute_type = "int8_float16" if device == "cuda" else "int8"

    print()
    print("[3/4] 加载 faster-whisper ...")
    print(f"模型        : {model_name}")
    print(f"设备        : {device}")
    print(f"compute_type: {compute_type}")
    print("首次运行可能会自动下载模型文件。")

    model = WhisperModel(
        model_name,
        device=device,
        compute_type=compute_type,
    )

    return model, compute_type


def transcribe_one(
    model,
    audio_path: Path,
    beam_size: int,
) -> tuple[str, str, float]:
    """
    强制按日语识别。

    这些 MP3 本身已经是单句角色语音，因此默认不启用 VAD，
    避免非常短的「うん」「え？」之类被误认为静音。
    """
    segments, info = model.transcribe(
        str(audio_path),
        language="ja",
        task="transcribe",
        beam_size=beam_size,
        vad_filter=False,
        condition_on_previous_text=False,
    )

    # segments 是 generator，真正推理发生在迭代时。
    texts: list[str] = []
    for segment in segments:
        piece = clean_asr_text(segment.text)
        if piece:
            texts.append(piece)

    ja_text = clean_asr_text("".join(texts))
    language = getattr(info, "language", "ja") or "ja"
    probability = float(
        getattr(info, "language_probability", 0.0) or 0.0
    )

    return ja_text, language, probability


def run_asr(
    records: list[dict],
    model_name: str,
    device: str,
    compute_type: str | None,
    beam_size: int,
    force_asr: bool,
) -> None:
    model, actual_compute_type = create_whisper_model(
        model_name=model_name,
        device=device,
        compute_type=compute_type,
    )

    print()
    print("[4/4] 开始识别日语台词...")

    total = len(records)

    for pos, record in enumerate(records, start=1):
        record = normalize_record(record)

        filename = record.get("filename", "")
        audio_path = AUDIO_DIR / filename

        if (
            not force_asr
            and record.get("ja_text", "").strip()
            and record.get("asr_status") == "ok"
        ):
            print(
                f"[{pos:03d}/{total:03d}] 跳过已识别: "
                f"{record['ja_text']}"
            )
            continue

        if not audio_path.exists() or audio_path.stat().st_size == 0:
            record["asr_status"] = "audio_missing"
            print(
                f"[{pos:03d}/{total:03d}] 缺少音频: {audio_path}"
            )
            save_metadata(records)
            continue

        print()
        print(f"[{pos:03d}/{total:03d}] ASR: {filename}")
        print(f"中文参考: {record.get('zh_text', '')}")

        try:
            ja_text, language, probability = transcribe_one(
                model=model,
                audio_path=audio_path,
                beam_size=beam_size,
            )

            record["ja_text"] = ja_text
            record["asr_model"] = model_name
            record["asr_language"] = language
            record["asr_language_probability"] = f"{probability:.4f}"
            record["asr_status"] = "ok" if ja_text else "empty"

            print(f"日文识别: {ja_text or '(空)'}")

        except Exception as exc:
            record["asr_model"] = model_name
            record["asr_status"] = f"error: {type(exc).__name__}"
            print(f"[识别失败] {exc}")

        # 每条立即落盘，支持断点续跑。
        save_metadata(records)

    print()
    print("处理完成。")
    print(f"音频目录: {AUDIO_DIR.resolve()}")
    print(f"CSV      : {CSV_PATH.resolve()}")
    print(f"JSON     : {JSON_PATH.resolve()}")
    print(f"ASR 模型 : {model_name}")
    print(f"设备      : {device}")
    print(f"计算类型  : {actual_compute_type}")


def main() -> None:
    args = parse_args()

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

    if args.asr_only:
        print("[1/4] --asr-only：跳过网页抓取与下载。")
        records = load_existing_metadata()
        print(f"[2/4] 已读取 {len(records)} 条 metadata。")
    else:
        session = make_session()
        records = prepare_records_from_web(session)

    run_asr(
        records=records,
        model_name=args.model,
        device=args.device,
        compute_type=args.compute_type,
        beam_size=args.beam_size,
        force_asr=args.force_asr,
    )


if __name__ == "__main__":
    main()
