test_single_pair.py 6.79 KB
#!/usr/bin/env python3
"""单对音频 DNA 匹配测试:将原版入库,再以翻唱查询,输出相似度结果。

用法:
    python scripts/aliyun_dna/test_single_pair.py \
        --ref "/path/to/原版.mp3" --ref-id original_id \
        --query "/path/to/翻唱.mp3" --query-id cover_id \
        --db-id <DNA库ID>  # 可选,默认读 .env
"""

import argparse
import json
import logging
import os
import sys
import time
import urllib.request
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
from dotenv import load_dotenv
load_dotenv(Path(__file__).resolve().parents[2] / ".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

logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
logger = logging.getLogger(__name__)

POLL_INTERVAL = 3
POLL_TIMEOUT = 300


def _get_oss_client() -> oss2.Bucket:
    auth = oss2.Auth(os.environ["ALIYUN_ICE_ACCESS_KEY_ID"], os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"])
    return oss2.Bucket(auth, f"https://{os.environ['ALIYUN_OSS_ENDPOINT']}", os.environ["ALIYUN_OSS_BUCKET"])


def _get_ice_client() -> ICEClient:
    region = os.environ.get("ALIYUN_ICE_REGION", "cn-hangzhou")
    config = OpenApiConfig(
        access_key_id=os.environ["ALIYUN_ICE_ACCESS_KEY_ID"],
        access_key_secret=os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"],
        endpoint=f"ice.{region}.aliyuncs.com",
        region_id=region,
    )
    return ICEClient(config)


def upload_to_oss(bucket: oss2.Bucket, audio_path: str, prefix: str, key: str) -> str:
    ext = Path(audio_path).suffix
    obj = f"{prefix}/{key}{ext}"
    size = Path(audio_path).stat().st_size
    logger.info("上传: %s (%.1fMB) -> %s", Path(audio_path).name, size / 1e6, obj)
    bucket.put_object_from_file(obj, audio_path)
    return f"oss://{bucket.bucket_name}/{obj}"


def submit_job(client: ICEClient, oss_url: str, primary_key: str, save_type: str) -> str:
    db_id = os.environ["ALIYUN_DNA_DB_ID"]
    req = ice_models.SubmitDNAJobRequest(
        input=ice_models.SubmitDNAJobRequestInput(type="OSS", media=oss_url),
        primary_key=primary_key,
        dbid=db_id,
        config=json.dumps({"SaveType": save_type, "MediaType": "audio"}),
    )
    resp = client.submit_dnajob(req)
    return resp.body.job_id


def poll_job(client: ICEClient, job_id: str) -> dict:
    t0 = time.time()
    while time.time() - t0 < POLL_TIMEOUT:
        resp = client.query_dnajob_list(ice_models.QueryDNAJobListRequest(job_ids=job_id))
        if resp.body.job_list:
            job = resp.body.job_list[0]
            if job.status == "Success":
                return {"status": "Success", "dna_result_url": job.dnaresult}
            elif job.status == "Fail":
                return {"status": "Fail", "error": job.message or "unknown"}
        time.sleep(POLL_INTERVAL)
    return {"status": "Timeout"}


def fetch_results(url: str) -> list[dict]:
    with urllib.request.urlopen(urllib.request.Request(url), timeout=30) as r:
        data = json.loads(r.read().decode())
    return data if isinstance(data, list) else [data]


def main():
    ap = argparse.ArgumentParser(description="单对音频 DNA 匹配测试")
    ap.add_argument("--ref", required=True, help="原版音频路径")
    ap.add_argument("--ref-id", required=True, help="原版歌曲 ID(入库主键)")
    ap.add_argument("--query", required=True, help="翻唱音频路径")
    ap.add_argument("--query-id", default="query_cover", help="查询 ID")
    ap.add_argument("--db-id", default="", help="覆盖 ALIYUN_DNA_DB_ID")
    ap.add_argument("--skip-upload-ref", action="store_true", help="原版已入库,跳过上传")
    ap.add_argument("--threshold", type=float, default=0.8, help="duplicate 判定阈值(默认 0.8)")
    args = ap.parse_args()

    if args.db_id:
        os.environ["ALIYUN_DNA_DB_ID"] = args.db_id

    for k in ["ALIYUN_ICE_ACCESS_KEY_ID", "ALIYUN_ICE_ACCESS_KEY_SECRET",
              "ALIYUN_OSS_BUCKET", "ALIYUN_OSS_ENDPOINT", "ALIYUN_DNA_DB_ID"]:
        if not os.environ.get(k):
            logger.error("缺少环境变量: %s", k)
            sys.exit(1)

    logger.info("DNA 库: %s", os.environ["ALIYUN_DNA_DB_ID"])
    bucket = _get_oss_client()
    client = _get_ice_client()

    # Step 1: 原版入库
    if not args.skip_upload_ref:
        logger.info("=== Step 1: 原版入库 (%s) ===", args.ref_id)
        ref_oss = upload_to_oss(bucket, args.ref, "dna-test-ref", args.ref_id)
        ref_job_id = submit_job(client, ref_oss, args.ref_id, "save")
        logger.info("入库作业已提交: job_id=%s,等待完成...", ref_job_id)
        ref_result = poll_job(client, ref_job_id)
        if ref_result["status"] != "Success":
            logger.error("原版入库失败: %s", ref_result)
            sys.exit(1)
        logger.info("原版入库成功,等待 5s 让索引生效...")
        time.sleep(5)
    else:
        logger.info("=== Step 1: 跳过原版入库(--skip-upload-ref)===")

    # Step 2: 翻唱查询
    logger.info("=== Step 2: 翻唱查询 (%s) ===", args.query_id)
    query_oss = upload_to_oss(bucket, args.query, "dna-test-query", args.query_id)
    query_job_id = submit_job(client, query_oss, args.query_id, "nosave")
    logger.info("查询作业已提交: job_id=%s,等待完成...", query_job_id)
    query_result = poll_job(client, query_job_id)

    if query_result["status"] != "Success":
        logger.error("查询失败: %s", query_result)
        sys.exit(1)

    # Step 3: 解析结果
    matches = fetch_results(query_result["dna_result_url"]) if query_result.get("dna_result_url") else []
    matches.sort(key=lambda x: x.get("GlobalSimilarity", 0), reverse=True)

    print("\n" + "=" * 60)
    print(f"查询: {Path(args.query).name}")
    print(f"原版 ID: {args.ref_id}  |  DNA库: {os.environ['ALIYUN_DNA_DB_ID']}")
    print("=" * 60)

    if not matches:
        print("无匹配结果")
    else:
        for i, m in enumerate(matches[:5], 1):
            pk = m.get("PrimaryKey", "")
            sim = m.get("GlobalSimilarity", 0)
            is_target = pk == args.ref_id
            mark = " <-- 目标原版" if is_target else ""
            print(f"  #{i} PrimaryKey={pk}  GlobalSimilarity={sim:.4f}{mark}")

        top1 = matches[0]
        top1_sim = top1.get("GlobalSimilarity", 0)
        top1_pk = top1.get("PrimaryKey", "")
        hit = top1_pk == args.ref_id
        is_dup = top1_sim >= args.threshold

        print()
        print(f"Top1 命中原版: {'是' if hit else '否'}")
        print(f"Top1 相似度:   {top1_sim:.4f}(阈值 {args.threshold})")
        print(f"判定为重复:    {'是 ✓' if is_dup else '否 ✗'}")

    print("=" * 60)


if __name__ == "__main__":
    main()