upload_dna_library.py 10.7 KB
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""批量上传音频到阿里云 OSS 并提交 DNA 入库作业。

流程:
1. 扫描音频文件
2. 降采样(可选,默认 16kHz 单声道)
3. 上传到 OSS
4. 调用 SubmitDNAJob 提交入库作业(SaveType=save, MediaType=audio)
5. 记录状态到 --state-file,支持断点续传

用法:
    # 入库(reference 原始歌曲)
conda activate hikoon-data-spider
python scripts/aliyun_dna/upload_dna_library.py \
    --audio-dir /Volumes/移动硬盘/composition_test \
    --state-file upload_dna_state.json

# 断点续传(跳过已成功入库的)
python scripts/aliyun_dna/upload_dna_library.py --retry-failed

# 不降采样(上传原始文件)
python scripts/aliyun_dna/upload_dna_library.py --ref-csv ... --no-resample
"""

import argparse
import csv
import json
import logging
import os
import sys
import tempfile
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent.parent))

from dotenv import load_dotenv
load_dotenv(Path(__file__).resolve().parent.parent.parent / ".env")

import oss2
from alibabacloud_ice20201109.client import Client as ICEClient
from alibabacloud_tea_openapi.models import Config as OpenApiConfig
from alibabacloud_ice20201109 import models as ice_models

logger = logging.getLogger(__name__)

SUPPORTED_EXTENSIONS = {".mp3", ".wav", ".flac", ".ogg", ".m4a", ".aac", ".wma"}
OSS_UPLOAD_PREFIX = "dna-audio"  # OSS 中的对象前缀
DEFAULT_SAMPLE_RATE = 16000  # 默认降采样到 16kHz


def _resample_audio(audio_path: str, target_sr: int = DEFAULT_SAMPLE_RATE) -> str | None:
    """降采样音频到目标采样率(单声道),返回临时文件路径。"""
    try:
        import librosa
        import soundfile as sf

        y, sr = librosa.load(audio_path, sr=target_sr, mono=True)
        tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
        sf.write(tmp.name, y, target_sr, subtype="PCM_16")
        tmp.close()

        orig_size = Path(audio_path).stat().st_size
        new_size = Path(tmp.name).stat().st_size
        ratio = (1 - new_size / orig_size) * 100 if orig_size > 0 else 0
        logger.info("  降采样: %dHz -> %dHz, 大小 %s -> %s (减少 %.0f%%)",
                    sr, target_sr,
                    _human_size(orig_size), _human_size(new_size), ratio)
        return tmp.name
    except ImportError:
        logger.warning("  缺少 librosa/soundfile,跳过降采样")
        return None
    except Exception as e:
        logger.warning("  降采样失败: %s, 使用原始文件", e)
        return None


def _human_size(n: int) -> str:
    for unit in ["B", "KB", "MB", "GB"]:
        if n < 1024:
            return f"{n:.1f}{unit}"
        n /= 1024
    return f"{n:.1f}TB"


def _get_oss_client() -> oss2.Bucket:
    """创建 OSS Bucket 客户端。"""
    access_key_id = os.environ["ALIYUN_ICE_ACCESS_KEY_ID"]
    access_key_secret = os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"]
    bucket_name = os.environ["ALIYUN_OSS_BUCKET"]
    endpoint = os.environ["ALIYUN_OSS_ENDPOINT"]

    auth = oss2.Auth(access_key_id, access_key_secret)
    return oss2.Bucket(auth, f"https://{endpoint}", bucket_name)


def _get_ice_client() -> ICEClient:
    """创建 ICE API 客户端。"""
    access_key_id = os.environ["ALIYUN_ICE_ACCESS_KEY_ID"]
    access_key_secret = os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"]
    region = os.environ.get("ALIYUN_ICE_REGION", "cn-hangzhou")

    config = OpenApiConfig(
        access_key_id=access_key_id,
        access_key_secret=access_key_secret,
        endpoint=f"ice.{region}.aliyuncs.com",
        region_id=region,
    )
    return ICEClient(config)


def upload_to_oss(oss_bucket: oss2.Bucket, audio_path: str, song_id: str) -> str:
    """上传音频文件到 OSS,返回 oss://bucket/object 地址。"""
    ext = Path(audio_path).suffix
    object_name = f"{OSS_UPLOAD_PREFIX}/{song_id}{ext}"

    size = Path(audio_path).stat().st_size
    logger.info("  上传到 OSS: %s (%s) -> %s", Path(audio_path).name, _human_size(size), object_name)
    oss_bucket.put_object_from_file(object_name, audio_path)
    return f"oss://{oss_bucket.bucket_name}/{object_name}"


def submit_dna_job(ice_client: ICEClient, oss_url: str, song_id: str, save_type: str = "save") -> str:
    """提交 DNA 作业,返回 JobId。

    save_type: save(入库)/ nosave(仅搜索不入库)
    """
    db_id = os.environ["ALIYUN_DNA_DB_ID"]

    config_json = json.dumps({
        "SaveType": save_type,
        "MediaType": "audio",
    })

    input_obj = ice_models.SubmitDNAJobRequestInput(
        type="OSS",
        media=oss_url,
    )

    request = ice_models.SubmitDNAJobRequest(
        input=input_obj,
        primary_key=song_id,
        dbid=db_id,
        config=config_json,
    )

    response = ice_client.submit_dnajob(request)
    job_id = response.body.job_id
    return job_id


def discover_audio_files(ref_csv: str | None, audio_dir: str | None) -> list[tuple[str, str]]:
    """返回 [(song_id, audio_path), ...]。

    优先用 ref_csv(reference.csv),否则扫描 audio_dir。
    """
    results = []

    if ref_csv:
        with open(ref_csv, newline="", encoding="utf-8") as f:
            reader = csv.DictReader(f)
            for row in reader:
                song_id = row["song_id"].strip()
                audio_path = row["audio_path"].strip()
                if audio_path and Path(audio_path).exists():
                    results.append((song_id, audio_path))
                else:
                    logger.warning("文件不存在: song_id=%s path=%s", song_id, audio_path)
    elif audio_dir:
        for audio_file in sorted(Path(audio_dir).rglob("*")):
            if audio_file.is_file() and audio_file.suffix.lower() in SUPPORTED_EXTENSIONS:
                if not audio_file.name.startswith("._"):
                    song_id = audio_file.stem.split("_")[0] if "_" in audio_file.stem else audio_file.stem
                    results.append((song_id, str(audio_file)))
    else:
        logger.error("请指定 --ref-csv 或 --audio-dir")
        sys.exit(1)

    return results


def load_state(state_file: str) -> dict:
    if not Path(state_file).exists():
        return {}
    with open(state_file, encoding="utf-8") as f:
        return json.load(f)


def save_state(state_file: str, state: dict) -> None:
    with open(state_file, "w", encoding="utf-8") as f:
        json.dump(state, f, ensure_ascii=False, indent=2)


def main() -> None:
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s [%(levelname)s] %(message)s",
    )

    parser = argparse.ArgumentParser(description="上传音频到 OSS 并提交 DNA 入库作业")
    parser.add_argument("--audio-dir", help="音频文件目录")
    parser.add_argument("--ref-csv", help="reference.csv 路径(优先使用)")
    parser.add_argument("--ext", default=".wav", help="音频文件扩展名(默认 .wav)")
    parser.add_argument("--state-file", default="upload_dna_state.json", help="上传状态文件")
    parser.add_argument("--retry-failed", action="store_true", help="重试上次失败的文件")
    parser.add_argument("--dry-run", action="store_true", help="只列出待上传文件,不实际操作")
    parser.add_argument("--save-type", default="save", choices=["save", "forcesave", "onlysave"],
                        help="DNA 存储类型(默认 save=去重入库)")
    parser.add_argument("--sample-rate", type=int, default=DEFAULT_SAMPLE_RATE,
                        help=f"降采样目标采样率(默认 {DEFAULT_SAMPLE_RATE}Hz)")
    parser.add_argument("--no-resample", action="store_true",
                        help="不降采样,上传原始文件")
    args = parser.parse_args()

    # 验证配置
    required_env = ["ALIYUN_ICE_ACCESS_KEY_ID", "ALIYUN_ICE_ACCESS_KEY_SECRET",
                    "ALIYUN_OSS_BUCKET", "ALIYUN_OSS_ENDPOINT", "ALIYUN_DNA_DB_ID"]
    for key in required_env:
        if not os.environ.get(key):
            logger.error("缺少环境变量: %s", key)
            sys.exit(1)

    audio_files = discover_audio_files(args.ref_csv, args.audio_dir)
    logger.info("发现 %d 个音频文件", len(audio_files))

    state = load_state(args.state_file)

    # 过滤已成功入库的文件
    pending = []
    for song_id, audio_path in audio_files:
        status = state.get(audio_path)
        if status == "ok":
            continue
        if status == "failed" and not args.retry_failed:
            continue
        pending.append((song_id, audio_path))

    ok_count = sum(1 for s in state.values() if s == "ok")
    fail_count_prev = sum(1 for s in state.values() if s == "failed")
    skipped_failed = fail_count_prev if not args.retry_failed else 0

    logger.info("待处理: %d 个(已跳过 %d 个成功 + %d 个失败)",
                len(pending), ok_count, skipped_failed)

    if args.dry_run:
        for song_id, audio_path in pending:
            print(f"  [dry-run] {audio_path} (song_id={song_id})")
        return

    oss_bucket = _get_oss_client()
    ice_client = _get_ice_client()

    success = 0
    fail = 0
    tmp_files = []  # 跟踪临时文件,结束时清理

    for i, (song_id, audio_path) in enumerate(pending, 1):
        logger.info("[%d/%d] 处理: %s (song_id=%s)", i, len(pending), Path(audio_path).name, song_id)

        upload_path = audio_path
        try:
            t0 = time.perf_counter()

            # 1. 降采样
            if not args.no_resample:
                resampled = _resample_audio(audio_path, target_sr=args.sample_rate)
                if resampled:
                    upload_path = resampled
                    tmp_files.append(resampled)

            # 2. 上传到 OSS
            oss_url = upload_to_oss(oss_bucket, upload_path, song_id)
            logger.info("  OSS 地址: %s", oss_url)

            # 3. 提交 DNA 入库作业
            job_id = submit_dna_job(ice_client, oss_url, song_id, save_type=args.save_type)
            elapsed_ms = (time.perf_counter() - t0) * 1000

            logger.info("  入库作业已提交: JobId=%s (%.0fms)", job_id, elapsed_ms)
            state[audio_path] = "ok"
            state[f"{audio_path}::job_id"] = job_id
            state[f"{audio_path}::oss_url"] = oss_url
            success += 1

        except Exception as e:
            logger.error("  处理失败: %s", e)
            state[audio_path] = "failed"
            state[f"{audio_path}::error"] = str(e)
            fail += 1

        save_state(args.state_file, state)

    # 清理临时文件
    for tmp in tmp_files:
        try:
            Path(tmp).unlink(missing_ok=True)
        except Exception:
            pass

    logger.info("入库完成: 成功 %d, 失败 %d,状态已保存到 %s", success, fail, args.state_file)


if __name__ == "__main__":
    main()