跳转至

tts-quality-eval

第6章 · Agent 的评估 · 配套项目 chapter6/tts-quality-eval

项目说明

实验 6-5:全自动 TTS 质量评估流水线

配套《深入理解 AI Agent》第 6 章「实验 6-5 ★★:构建全自动 TTS 质量评估流水线」。

用多个 TTS provider / 配置(OpenAI、ElevenLabs、Fish Audio、Minimax、豆包,或同 一家的不同 model / voice / speed)合成同一组带挑战性的参考文本,再用 多模态 LLM-as-a-Judge 的思路对合成语音按 Rubric 逐维度打分,最后汇总成一张 对比表,反映不同 provider / 配置在准确性 / 自然度上的优劣。

目的

回答工程中的实际问题:同一段文本,tts-1tts-1-hd 有多大差距?换 voice、把 语速调到 1.5x 会牺牲多少质量? 本 demo 把这类对比做成一条命令跑通、可复现的流水线。

评审维度与 Rubric

对每条合成语音,先测出客观特征(时长、语速、字错误率),再让评审模型按 1–5 分打分:

维度 含义
清晰度 转写是否与原文一致(漏字/错字/多字越多分越低,对应准确性维度)
自然度 语速是否接近自然朗读(中文约 4–6 字/秒,过快/过慢都扣分)
停顿节奏 结合语速与文本长度判断节奏是否合理(过快常意味吞字)
整体 综合印象分

客观指标 CER(字错误率)/ 字准确率:把 Whisper 回译文本与原文归一化(去标点空白、 统一大小写)后做字符级编辑距离,CER = 编辑距离 / 参考字数字准确率 = 1 - CER。 中文按字级计算(等价于书中 WER 的可懂度维度)。

Provider 适配说明

  • TTS 合成(多 provider):对应书中「接入主流服务:OpenAI、ElevenLabs、Fish Audio、 Minimax、豆包」。每个 provider 按各家公开 REST 接口实现(OpenAI 走官方 SDK,其余走内置 urllib,无额外依赖)。默认(不加 --providers)只跑 OpenAI 的 4 个配置,保证单个 OPENAI_API_KEY 即可零配置跑通;--providers openai,minimax,... 做跨服务商横向对比。 各 provider 所需环境变量与 voice 字段语义见 python demo.py --list-providers
provider 环境变量 voice 语义
openai OPENAI_API_KEY alloy/nova…;model=tts-1 / tts-1-hd / gpt-4o-mini-tts
elevenlabs ELEVENLABS_API_KEY voice_id;model 默认 eleven_multilingual_v2
fishaudio FISH_API_KEY(别名 FISHAUDIO_API_KEY reference_id(留空用默认音色)
minimax MINIMAX_API_KEY + MINIMAX_GROUP_ID voice_id;model 默认 speech-01-turbo
doubao DOUBAO_APP_ID + DOUBAO_ACCESS_TOKEN voice_type(火山引擎)

说明:本仓库仅 OpenAI 路径经端到端验证;其余四家按各自公开 REST 文档实现,请用自己 账号可用的 voice/model 覆盖 config.PROVIDER_CONFIGS 后使用。缺对应 key 时该 provider 的行会被记为失败,不中断整表。 - 质量评审(默认):用 Whisper(whisper-1)把合成语音回译成文本算 CER,再用 gpt-5.6-luna(当前廉价旗舰)基于「转写文本 + 时长 + 语速 + CER」按 Rubric 打分。 转写时用简体中文提示语引导 Whisper 输出简体,避免繁体字形差异虚高 CER。 凭据/回退:TTS 合成与 Whisper 回译必须走 OpenAI 直连OPENAI_API_KEY, OpenRouter 不提供音频/转写);仅 LLM Rubric 的 chat 评审支持 OpenRouter 回退—— gpt-5.x 直连需组织实名认证,故只要设置了 OPENROUTER_API_KEY,评审就优先走 OpenRouter(gpt-* 映射为 openai/*)。 - 质量评审(可选,书中方案)--geminiGemini 多模态直接「听」音频打分 (原文 + 音频 + Rubric 一起输入),需 GEMINI_API_KEY。默认模型为 gemini-3.5-flash(已验证支持音频输入);代码会先探测 /models,若该名不可用再 自动回退到当前可用模型(如 gemini-2.5-pro)。

书中用 Gemini 直接听合成语音打分(本 demo 默认 gemini-3.5-flash,已验证支持音频); 默认改用「Whisper 回译 + LLM Rubric」以便零额外配置即可跑通,同时保留 --gemini 开关复现书中方案。两者的 区别:Gemini 能直接感知音色/韵律/情感;回译方案只能基于可测特征做保守推断(见「局限」)。

文件

文件 说明
config.py 模型名与单价、provider 注册表(PROVIDERS / PROVIDER_CONFIGS)、TTS 配置集合、测试语料
pipeline.py 多 provider 合成分发 / ffprobe 时长 / Whisper 回译 / CER 计算 / LLM Rubric / 可选 Gemini
demo.py 入口:多配置 × 多语料跑全流程,打印逐条明细 + 对比汇总表
requirements.txt / env.example 依赖与环境变量示例

运行

pip install -r requirements.txt          # 只需 openai
brew install ffmpeg                        # 提供 ffprobe(时长探测)
export OPENAI_API_KEY=sk-...

python demo.py            # 默认:4 个 OpenAI 配置 × 4 条语料,Whisper 回译 + LLM Rubric
python demo.py --quick   # 只用前 2 条语料,快速冒烟
python demo.py --extra   # 额外加入 gpt-4o-mini-tts 配置
python demo.py --gemini  # 评审改用 Gemini 多模态直接听音频(需 GEMINI_API_KEY)
python demo.py --fresh   # 忽略已有音频全部重合成

# 多 provider / 自定义输入(新增)
python demo.py --providers openai,minimax,elevenlabs   # 跨服务商横向对比(需各自 key)
python demo.py --text '2026年营收增长37.5%'             # 用一段自定义文本替换语料库
python demo.py --judge-model gpt-5.6-luna                 # 覆盖 LLM 评审模型
python demo.py --output ./runs/exp1                     # 自定义输出目录

# 离线(无需任何 API key)
python demo.py --list-providers   # 查看所有 provider 及配置状态
python demo.py --dump-rubric      # 查看 Rubric 维度定义

完整参数见 python demo.py --help(全中文)。合成音频写入 output/(已被 .gitignore 忽略),结构化结果写入 output/results.json(可用 --output 改目录)。 幂等:默认复用已存在的音频,重复运行不会重复合成。

测试语料

4 条覆盖不同挑战点:数字/百分比/日期、多音字(行/长/重/还)、长句新闻文体、 专有名词 + 感叹情感。可在 config.pyCORPUS 中增删。

健壮性

  • OPENAI_API_KEY 立即清晰报错退出;缺 ffprobe 给出安装提示。
  • 单个(配置, 语料)在合成/转写/评审任一步失败,只把该条记为失败,不中断整表, 汇总表按成功条数聚合。
  • OpenAI 客户端带自动重试(max_retries=5)缓解偶发网络抖动。
  • ffprobe 调用检查返回码与输出可解析性。

局限

  • 默认评审看不到音频本身:只基于回译文本 + 客观特征推断,无法直接判断音色一致度、 真实韵律与情感表达(书中的 Gemini 方案能,用 --gemini 复现)。因此「自然度/情感」 维度是保守估计。音色一致性维度需参考语音,本 demo 未覆盖。
  • CER 依赖 Whisper 转写质量,Whisper 自身错误会引入噪声;数字/专名可能因书写形式 (阿拉伯数字 vs 中文数字)产生非发音性差异。
  • Rubric 由 LLM 打分,存在评审模型偏好;分数用于相对对比而非绝对基准。

源代码

config.py

"""实验 6-5:全自动 TTS 质量评估流水线 —— 配置与测试语料。

本模块集中管理:
  - 用到的 OpenAI 模型名与计费单价(仅供参考成本估算);
  - 多个 TTS「配置」(model / voice / speed 的组合,作为待对比的对象);
  - 一组带挑战性的参考文本(数字 / 多音字 / 长句 / 专有名词 + 情感)。
"""

from dataclasses import dataclass, field

# ---------------------------------------------------------------------------
# 模型名(均为 OpenAI,读 OPENAI_API_KEY)。
# ---------------------------------------------------------------------------
WHISPER_MODEL = "whisper-1"        # 语音转写(回译),用于计算 WER/字准确率(须走 OpenAI 直连)
JUDGE_MODEL = "gpt-5.6-luna"       # LLM Rubric 评审模型(当前廉价旗舰;chat 调用可回退 OpenRouter)

# 可选的 Gemini 音频评审(书中方案)。默认用当前廉价旗舰 gemini-3.5-flash(已验证支持
# 音频输入,能直接「听」合成语音)。模型名可能随时间过期,运行时会通过 REST /models
# 探测校正。仅当 --gemini 开启时才会用到。
GEMINI_MODEL_DEFAULT = "gemini-3.5-flash"

# 计费单价(美元),仅用于打印粗略成本,不影响评分。数值随官方调整可能变化。
PRICE = {
    "tts-1": 15.0 / 1_000_000,          # $ / 字符
    "tts-1-hd": 30.0 / 1_000_000,       # $ / 字符
    "gpt-4o-mini-tts": 12.0 / 1_000_000,
    "whisper-1": 0.006 / 60,            # $ / 秒
}


@dataclass
class TTSConfig:
    """一个待评估的 TTS 配置。name 需在整表内唯一。

    provider 指明合成走哪个服务商(openai / elevenlabs / fishaudio / minimax /
    doubao)。model / voice / speed 的语义由各 provider 自行解释:例如 elevenlabs
    的 voice 是 voice_id,fishaudio 的 voice 是 reference_id(可留空用默认音色)。
    """

    name: str
    model: str
    voice: str
    speed: float = 1.0
    provider: str = "openai"

    def supports_speed(self) -> bool:
        # 只有部分 provider/模型支持 speed 参数;不支持时忽略该字段。
        if self.provider == "openai":
            return self.model in ("tts-1", "tts-1-hd")
        return self.provider in ("minimax", "doubao")


# ---------------------------------------------------------------------------
# 多 provider 注册表(对应书中「接入主流服务:OpenAI、ElevenLabs、Fish Audio、
# Minimax、豆包」)。每个 provider 声明所需环境变量与一个代表性配置,便于跨服务商
# 横向对比。除 OpenAI 外均按各家公开 REST 接口实现,缺 key 时该 provider 的行会被
# 记为失败而不影响整表(见 demo.py)。
# ---------------------------------------------------------------------------
# 环境变量别名:同一凭据可能有多个历史/惯用名,任意一个被设置即视为已配置。
ENV_ALIASES = {
    "FISH_API_KEY": ("FISH_API_KEY", "FISHAUDIO_API_KEY"),
}


def env_get(name: str) -> str:
    """读取环境变量,支持 ENV_ALIASES 中登记的别名,返回第一个非空值(已 strip)。"""
    import os
    for n in ENV_ALIASES.get(name, (name,)):
        val = os.environ.get(n, "").strip()
        if val:
            return val
    return ""


@dataclass
class ProviderInfo:
    key: str                # 内部标识(--providers 用)
    label: str              # 展示名
    env: tuple              # 该 provider 合成所需的环境变量名
    note: str               # 一句话说明 voice 字段语义等

    def configured(self) -> bool:
        return all(env_get(e) for e in self.env)


PROVIDERS = {
    "openai": ProviderInfo(
        "openai", "OpenAI", ("OPENAI_API_KEY",),
        "voice=alloy/nova/…,model=tts-1/tts-1-hd/gpt-4o-mini-tts;本仓库唯一端到端验证过的 provider。",
    ),
    "elevenlabs": ProviderInfo(
        "elevenlabs", "ElevenLabs", ("ELEVENLABS_API_KEY",),
        "voice=voice_id,model 默认 eleven_multilingual_v2(多语言/中文)。",
    ),
    "fishaudio": ProviderInfo(
        "fishaudio", "Fish Audio", ("FISH_API_KEY",),
        "voice=reference_id(留空用默认音色),走 /v1/tts;key 亦可用别名 FISHAUDIO_API_KEY。",
    ),
    "minimax": ProviderInfo(
        "minimax", "Minimax", ("MINIMAX_API_KEY", "MINIMAX_GROUP_ID"),
        "voice=voice_id,model 默认 speech-01-turbo;需额外 GroupId。",
    ),
    "doubao": ProviderInfo(
        "doubao", "豆包(火山引擎)", ("DOUBAO_APP_ID", "DOUBAO_ACCESS_TOKEN"),
        "voice=voice_type,走 openspeech.bytedance.com;鉴权头为 'Bearer;{token}'。",
    ),
}

# 各 provider 的代表性配置(--providers 选中时,每个 provider 取这一条参与对比)。
# 非 OpenAI 的 voice/model 取各家常见默认值,可在此按账号可用音色调整。
PROVIDER_CONFIGS = {
    "openai": TTSConfig("openai-alloy", provider="openai", model="tts-1", voice="alloy"),
    "elevenlabs": TTSConfig("elevenlabs-multi", provider="elevenlabs",
                            model="eleven_multilingual_v2", voice="21m00Tcm4TlvDq8ikWAM"),
    "fishaudio": TTSConfig("fishaudio-default", provider="fishaudio",
                          model="speech-1.5", voice=""),
    "minimax": TTSConfig("minimax-turbo", provider="minimax",
                        model="speech-01-turbo", voice="male-qn-qingse"),
    "doubao": TTSConfig("doubao-tts", provider="doubao",
                       model="volcano_tts", voice="zh_female_qingxin"),
}


# 默认对比的配置集合:覆盖 model(tts-1 vs tts-1-hd)、voice、speed 三个维度,
# 便于观察不同配置在准确性/自然度上的差异。默认全部走 OpenAI 以保证零额外配置跑通。
TTS_CONFIGS = [
    TTSConfig("tts1-alloy-1.0", model="tts-1", voice="alloy", speed=1.0),
    TTSConfig("tts1hd-alloy-1.0", model="tts-1-hd", voice="alloy", speed=1.0),
    TTSConfig("tts1-nova-1.0", model="tts-1", voice="nova", speed=1.0),
    TTSConfig("tts1-alloy-1.5", model="tts-1", voice="alloy", speed=1.5),
]

# 可选加入(--extra 开启):gpt-4o-mini-tts。默认不加入以保证一定跑通。
EXTRA_CONFIGS = [
    TTSConfig("4omini-nova-1.0", model="gpt-4o-mini-tts", voice="nova", speed=1.0),
]


@dataclass
class Sample:
    """一条参考文本 + 期望情感标签(供 Rubric 情感维度参考)。"""

    id: str
    text: str
    challenge: str      # 该样本主要考察的挑战点
    emotion: str = "中性"


# 多样化测试语料:数字/日期、多音字、长句、专有名词+情感。
CORPUS = [
    Sample(
        id="num",
        text="2026年第三季度营收增长了37.5%,同比提升12个百分点。",
        challenge="数字/百分比/日期",
        emotion="中性",
    ),
    Sample(
        id="polyphone",
        text="银行行长正在重新调整这件事的重点,长此以往,还得还清所有欠款。",
        challenge="多音字(行/长/重/还)",
        emotion="中性",
    ),
    Sample(
        id="long",
        text="据报道,随着人工智能技术的快速发展,越来越多的企业开始将大语言模型"
             "应用于客户服务、内容创作和数据分析等场景,从而显著提升了运营效率。",
        challenge="长句/新闻文体",
        emotion="中性",
    ),
    Sample(
        id="emotion",
        text="太棒了!OpenAI 刚刚发布的新模型在 GAIA 基准测试上表现惊人!",
        challenge="专有名词 + 感叹情感",
        emotion="兴奋",
    ),
]

demo.py

"""实验 6-5:全自动 TTS 质量评估流水线 —— 一条命令跑通。

    python demo.py                      # 默认 4 个 OpenAI 配置 x 4 条语料
    python demo.py --providers openai,minimax   # 跨服务商横向对比
    python demo.py --text '一段话'       # 自定义文本
    python demo.py --gemini             # 评审改用 Gemini 多模态直接听音频(需 GEMINI_API_KEY)
    python demo.py --quick              # 只用前 2 条语料,快速冒烟
    python demo.py --list-providers     # 离线:查看 provider 及配置状态
    python demo.py --dump-rubric        # 离线:查看 Rubric 维度定义

流程:多 provider TTS 合成 -> ffprobe 时长 -> Whisper 回译 -> CER/字准确率
      -> LLM/Gemini Rubric 打分 -> 打印逐条明细 + 配置对比汇总表。
幂等:音频写入 output/ 并复用(除非 --fresh)。完整参数见 `python demo.py --help`。
"""

import argparse
import json
import os
import sys
import time
import traceback
from statistics import mean

import config
import pipeline

OUT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "output")


def load_env():
    """极简 .env 加载(不引第三方依赖)。"""
    path = os.path.join(os.path.dirname(os.path.abspath(__file__)), ".env")
    if os.path.exists(path):
        for line in open(path, encoding="utf-8"):
            line = line.strip()
            if line and not line.startswith("#") and "=" in line:
                k, v = line.split("=", 1)
                os.environ.setdefault(k.strip(), v.strip())


def audio_path(cfg_name: str, sample_id: str) -> str:
    return os.path.join(OUT_DIR, f"{cfg_name}__{sample_id}.mp3")


def evaluate_one(cfg, sample, use_gemini: bool, fresh: bool,
                 judge_model: str = None) -> dict:
    """对单个 (配置, 语料) 跑完整链路。任一步失败返回 error 记录,不抛出。"""
    rec = {"config": cfg.name, "sample": sample.id, "challenge": sample.challenge,
           "provider": getattr(cfg, "provider", "openai"), "ok": False, "error": None}
    path = audio_path(cfg.name, sample.id)
    try:
        # 1) 合成(幂等:已存在且非 fresh 则复用)
        if fresh or not os.path.exists(path) or os.path.getsize(path) == 0:
            pipeline.synthesize(cfg, sample.text, path)
        # 2) 时长
        dur = pipeline.probe_duration(path)
        # 3) 回译
        hyp = pipeline.transcribe(path)
        # 4) CER / 字准确率
        er = pipeline.char_error_rate(sample.text, hyp)
        # 5) Rubric 打分
        if use_gemini:
            rub = pipeline.judge_gemini_audio(sample.text, sample.emotion, path)
        else:
            rub = pipeline.judge_rubric(sample.text, sample.emotion, hyp, dur, er.cer,
                                        model=judge_model)
        rec.update({
            "ok": True, "duration": dur, "hypothesis": hyp,
            "cer": er.cer, "accuracy": er.accuracy,
            "speed": (er.ref_len / dur) if dur else 0.0,
            "scores": rub.scores, "reasons": rub.reasons,
        })
    except Exception as e:  # 单条失败不影响整表
        rec["error"] = f"{type(e).__name__}: {e}"
    return rec


def fmt(x, nd=2):
    return f"{x:.{nd}f}" if isinstance(x, (int, float)) else str(x)


def print_detail(rec, sample_text):
    head = f"[{rec['config']} | {rec['sample']}] {rec['challenge']}"
    if not rec["ok"]:
        print(f"  {head}\n    !! 失败: {rec['error']}")
        return
    print(f"  {head}")
    print(f"    原文  : {sample_text}")
    print(f"    回译  : {rec['hypothesis']}")
    print(f"    时长  : {fmt(rec['duration'])}s   语速: {fmt(rec['speed'])} 字/秒"
          f"   CER: {fmt(rec['cer'],3)}   字准确率: {fmt(rec['accuracy']*100,1)}%")
    s, r = rec["scores"], rec["reasons"]
    for dim in pipeline.RUBRIC_DIMENSIONS:
        print(f"    {dim:<4}: {s.get(dim,'-')}/5  {r.get(dim,'')}")


def summarize(records):
    """按配置聚合:平均 CER、平均字准确率、各 Rubric 维度均分、成功数。"""
    by_cfg = {}
    for rec in records:
        by_cfg.setdefault(rec["config"], []).append(rec)
    rows = []
    for cfg_name, recs in by_cfg.items():
        ok = [r for r in recs if r["ok"]]
        row = {"config": cfg_name, "n_ok": len(ok), "n": len(recs)}
        if ok:
            row["cer"] = mean(r["cer"] for r in ok)
            row["accuracy"] = mean(r["accuracy"] for r in ok)
            for dim in pipeline.RUBRIC_DIMENSIONS:
                row[dim] = mean(r["scores"].get(dim, 0) for r in ok)
        rows.append(row)
    # 按整体分降序、CER 升序排序
    rows.sort(key=lambda x: (-x.get("整体", 0), x.get("cer", 1)))
    return rows


def print_table(rows):
    cols = ["整体", "清晰度", "自然度", "停顿节奏"]
    header = (f"{'配置':<18}{'成功':>6}{'字准确率':>10}{'CER':>8}"
              + "".join(f"{c:>9}" for c in cols))
    print(header)
    print("-" * 74)
    for r in rows:
        ok_str = f"{r['n_ok']}/{r['n']}"
        if not r.get("n_ok"):
            print(f"{r['config']:<18}{ok_str:>6}   (全部失败)")
            continue
        acc = f"{r['accuracy']*100:.1f}%"
        line = f"{r['config']:<18}{ok_str:>6}{acc:>10}{r['cer']:>8.3f}"
        line += "".join(f"{r.get(c,0):>9.2f}" for c in cols)
        print(line)


def print_providers():
    """离线打印所有可用 TTS provider 及其配置状态(无需任何 API key)。"""
    print("可用 TTS provider(书中:OpenAI / ElevenLabs / Fish Audio / Minimax / 豆包):\n")
    for key, p in config.PROVIDERS.items():
        state = "已配置" if p.configured() else "未配置"
        env = " + ".join(p.env)
        print(f"  [{key}]  {p.label}   ({state};需 {env})")
        print(f"      {p.note}")
    print("\n用 --providers openai,minimax 选择跨服务商横向对比(默认仅 OpenAI)。")
    print("非 OpenAI provider 需各自的 key(见 env.example);缺 key 时该行记为失败,不中断整表。")


def print_rubric():
    """离线打印 Rubric 维度定义(无需任何 API key)。"""
    print("TTS 质量评估 Rubric(1-5 分,5 最好):\n")
    for dim in pipeline.RUBRIC_DIMENSIONS:
        print(f"  {dim}{pipeline.RUBRIC_DESCRIPTIONS.get(dim, '')}")
    print("\n默认(Whisper 回译 + LLM)评审基于「转写文本 + 时长 + 语速 + CER」保守打分;")
    print("--gemini 让多模态模型直接听音频,可覆盖书中「情感表达 / 音色一致性」维度。")


def main():
    global OUT_DIR
    ap = argparse.ArgumentParser(
        description="全自动 TTS 质量评估流水线(实验 6-5):多 provider 合成 + 多模态 LLM Rubric 评审",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog="示例:\n"
               "  python demo.py                          默认 4 个 OpenAI 配置 × 4 条语料\n"
               "  python demo.py --providers openai,minimax   跨服务商横向对比\n"
               "  python demo.py --text '今天天气不错' --gemini   自定义文本 + Gemini 多模态评审\n"
               "  python demo.py --list-providers         离线查看 provider 及配置状态\n"
               "  python demo.py --dump-rubric            离线查看 Rubric 维度定义",
    )
    ap.add_argument("--text", metavar="文本",
                    help="用一段自定义文本替换测试语料库(只评这一句)")
    ap.add_argument("--providers", metavar="列表",
                    help="逗号分隔的 provider(openai,elevenlabs,fishaudio,minimax,doubao),"
                         "每个取代表性配置做横向对比;默认仅 OpenAI 的多配置")
    ap.add_argument("--judge-model", metavar="模型", dest="judge_model",
                    help=f"覆盖 LLM 评审模型(默认 {config.JUDGE_MODEL});--gemini 时不生效")
    ap.add_argument("--output", metavar="目录",
                    help=f"输出目录(音频 + results.json),默认 {OUT_DIR}")
    ap.add_argument("--extra", action="store_true", help="额外加入 gpt-4o-mini-tts 配置")
    ap.add_argument("--gemini", action="store_true", help="用 Gemini 多模态直接听音频评审(需 GEMINI_API_KEY)")
    ap.add_argument("--quick", action="store_true", help="只用前 2 条语料快速冒烟")
    ap.add_argument("--fresh", action="store_true", help="忽略已有音频,全部重新合成")
    ap.add_argument("--list-providers", action="store_true", dest="list_providers",
                    help="离线打印所有 TTS provider 及配置状态后退出(无需 key)")
    ap.add_argument("--dump-rubric", action="store_true", dest="dump_rubric",
                    help="离线打印 Rubric 维度定义后退出(无需 key)")
    args = ap.parse_args()

    load_env()

    # 离线路径:不联网、不需要任何 key,打印后直接退出。
    if args.list_providers:
        print_providers()
        return
    if args.dump_rubric:
        print_rubric()
        return

    if args.output:
        OUT_DIR = os.path.abspath(args.output)
    os.makedirs(OUT_DIR, exist_ok=True)

    if not os.environ.get("OPENAI_API_KEY", "").strip():
        print("错误:缺少 OPENAI_API_KEY(回译/默认评审需要)。请 export 或写入 .env 后重试。",
              file=sys.stderr)
        sys.exit(1)

    # 选择待对比的配置:--providers 优先(跨服务商),否则默认 OpenAI 多配置。
    if args.providers:
        configs = []
        for key in [p.strip() for p in args.providers.split(",") if p.strip()]:
            if key not in config.PROVIDER_CONFIGS:
                print(f"错误:未知 provider {key!r}。可用:{', '.join(config.PROVIDER_CONFIGS)}",
                      file=sys.stderr)
                sys.exit(1)
            configs.append(config.PROVIDER_CONFIGS[key])
    else:
        configs = list(config.TTS_CONFIGS)
        if args.extra:
            configs += config.EXTRA_CONFIGS

    if args.text:
        corpus = [config.Sample(id="custom", text=args.text,
                                challenge="自定义文本", emotion="中性")]
    else:
        corpus = config.CORPUS[:2] if args.quick else config.CORPUS

    judge_model = args.judge_model or config.JUDGE_MODEL
    mode = ("Gemini 多模态音频评审" if args.gemini
            else f"Whisper 回译 + LLM Rubric({judge_model})")
    providers_used = sorted({getattr(c, "provider", "openai") for c in configs})
    print("=" * 72)
    print(f"实验 6-5:全自动 TTS 质量评估流水线")
    print(f"评审模式:{mode}")
    print(f"参与 provider:{', '.join(providers_used)}")
    print(f"配置数:{len(configs)}   语料数:{len(corpus)}   "
          f"共 {len(configs)*len(corpus)} 条待评估")
    print("=" * 72)

    records = []
    t0 = time.time()
    for cfg in configs:
        print(f"\n### 配置 {cfg.name}  (provider={getattr(cfg,'provider','openai')}, "
              f"model={cfg.model}, voice={cfg.voice}, speed={cfg.speed})")
        for sample in corpus:
            rec = evaluate_one(cfg, sample, args.gemini, args.fresh,
                               judge_model=None if args.gemini else args.judge_model)
            print_detail(rec, sample.text)
            records.append(rec)

    rows = summarize(records)
    print("\n" + "=" * 72)
    print("配置对比汇总(按 整体 分降序)")
    print("=" * 72)
    print_table(rows)

    ok = sum(1 for r in records if r["ok"])
    print(f"\n完成:{ok}/{len(records)} 条成功,耗时 {time.time()-t0:.1f}s。")

    # 落盘结构化结果,便于二次分析
    out_json = os.path.join(OUT_DIR, "results.json")
    with open(out_json, "w", encoding="utf-8") as f:
        json.dump({"records": records, "summary": rows}, f,
                  ensure_ascii=False, indent=2)
    print(f"明细结果已写入 {out_json}")


if __name__ == "__main__":
    main()

pipeline.py

"""TTS 质量评估流水线的核心步骤。

一条评估链路:
  合成(OpenAI TTS) -> 时长探测(ffprobe) -> 回译(Whisper) -> 计算 CER/字准确率
      -> LLM Rubric 打分(gpt-5.6-luna) [可选: Gemini 音频评审 gemini-3.5-flash]

说明:TTS 合成与 Whisper 回译必须走 OpenAI 直连(OpenRouter 不提供音频/转写);
仅 LLM Rubric 的 chat 评审支持 OpenRouter 回退——gpt-5.x 直连需组织实名认证,
故只要有 OPENROUTER_API_KEY 就优先经 OpenRouter 调评审模型(见 get_judge_client_and_model)。

所有对外函数都做了健壮性处理:单条失败抛出带上下文的异常,由 demo.py 捕获后
在汇总表里记为失败,而不会中断整表。
"""

import base64
import json
import os
import re
import shutil
import subprocess
from dataclasses import dataclass, field
from typing import Optional

from openai import OpenAI

import config

# ---------------------------------------------------------------------------
# 客户端(带自动重试,缓解偶发的网络抖动)。
# ---------------------------------------------------------------------------
_client: Optional[OpenAI] = None


def get_client() -> OpenAI:
    """OpenAI 直连 client:用于 TTS 合成与 Whisper 回译(这两项不能走 OpenRouter)。"""
    global _client
    if _client is None:
        key = os.environ.get("OPENAI_API_KEY", "").strip()
        if not key:
            raise RuntimeError(
                "缺少 OPENAI_API_KEY(TTS 合成 / Whisper 回译需 OpenAI 直连)。"
                "请 `export OPENAI_API_KEY=sk-...` 或写入 .env。"
            )
        _client = OpenAI(api_key=key, max_retries=5, timeout=60.0)
    return _client


# ---------------------------------------------------------------------------
# LLM Rubric 评审客户端:支持 OpenRouter 回退。
# gpt-5.x 直连 OpenAI 需组织实名认证,只要有 OPENROUTER_API_KEY 就优先走 OpenRouter。
# 注意:仅 chat 评审可回退;TTS / Whisper 仍需 OpenAI 直连(见 get_client)。
# ---------------------------------------------------------------------------
OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
_judge_client: Optional[OpenAI] = None
_judge_client_kind: str = ""


def _to_openrouter_model(model: str) -> str:
    """把模型名映射成 OpenRouter id:含 '/' 视为原生 id;gpt-* -> openai/*;
    claude-* -> anthropic/claude-opus-4.8;其余回退到 openai/gpt-5.6-luna。"""
    if "/" in model:
        return model
    if model.startswith("gpt-"):
        return "openai/" + model
    if model.startswith("claude-"):
        return "anthropic/claude-opus-4.8"
    return "openai/gpt-5.6-luna"


def get_judge_client_and_model(model: str):
    """构造 LLM 评审用的 client 并返回 (client, 实际模型名)。

    回退:gpt-5.x 且有 OPENROUTER_API_KEY -> 优先 OpenRouter;否则有 OPENAI_API_KEY ->
    直连;否则有 OPENROUTER_API_KEY -> OpenRouter(模型名映射);皆无 -> 清晰报错。
    """
    global _judge_client, _judge_client_kind
    primary = os.environ.get("OPENAI_API_KEY", "").strip()
    orkey = os.environ.get("OPENROUTER_API_KEY", "").strip()
    prefer_or = bool(orkey) and model.startswith("gpt-5")

    if not prefer_or and primary:
        if _judge_client_kind != "openai":
            _judge_client = OpenAI(api_key=primary, max_retries=5, timeout=60.0)
            _judge_client_kind = "openai"
        return _judge_client, model
    if orkey:
        if _judge_client_kind != "openrouter":
            _judge_client = OpenAI(base_url=OPENROUTER_BASE_URL, api_key=orkey,
                                   max_retries=5, timeout=60.0)
            _judge_client_kind = "openrouter"
        return _judge_client, _to_openrouter_model(model)
    if primary:
        if _judge_client_kind != "openai":
            _judge_client = OpenAI(api_key=primary, max_retries=5, timeout=60.0)
            _judge_client_kind = "openai"
        return _judge_client, model
    raise RuntimeError(
        "缺少 OPENAI_API_KEY / OPENROUTER_API_KEY,无法运行 LLM Rubric 评审。"
    )


# ---------------------------------------------------------------------------
# 1) TTS 合成(多 provider 分发)
# ---------------------------------------------------------------------------
def synthesize(cfg: config.TTSConfig, text: str, out_path: str) -> None:
    """按 cfg.provider 分发到对应服务商合成语音,写入 out_path(mp3)。失败抛异常。

    OpenAI 走官方 SDK;其余服务商按各家公开 REST 接口用内置 urllib 调用,
    不引入额外依赖。缺少对应 key 时抛出带上下文的异常,由上层记为该行失败。
    """
    fn = _SYNTH_DISPATCH.get(cfg.provider)
    if fn is None:
        raise RuntimeError(
            f"未知 provider: {cfg.provider!r}(可选:{', '.join(_SYNTH_DISPATCH)})"
        )
    audio = fn(cfg, text)
    if not audio:
        raise RuntimeError(f"{cfg.provider} TTS 返回空音频")
    with open(out_path, "wb") as f:
        f.write(audio)


def _require_env(name: str) -> str:
    # 走 config.env_get 以支持环境变量别名(如 Fish 的 FISH_API_KEY / FISHAUDIO_API_KEY)。
    val = config.env_get(name)
    if not val:
        raise RuntimeError(f"缺少环境变量 {name},无法用该 provider 合成。")
    return val


def _http_post(url: str, body: dict, headers: dict, timeout: float = 90.0) -> bytes:
    """POST JSON,返回原始响应字节。非 2xx 抛出带响应体片段的异常。"""
    import urllib.error
    import urllib.request
    req = urllib.request.Request(
        url, data=json.dumps(body).encode(),
        headers={"Content-Type": "application/json", **headers}, method="POST",
    )
    try:
        with urllib.request.urlopen(req, timeout=timeout) as r:
            return r.read()
    except urllib.error.HTTPError as e:
        detail = e.read().decode("utf-8", "replace")[:300]
        raise RuntimeError(f"HTTP {e.code}: {detail}") from None


def _synth_openai(cfg: config.TTSConfig, text: str) -> bytes:
    kwargs = dict(model=cfg.model, voice=cfg.voice, input=text)
    if cfg.supports_speed() and abs(cfg.speed - 1.0) > 1e-6:
        kwargs["speed"] = cfg.speed
    return get_client().audio.speech.create(**kwargs).content


def _synth_elevenlabs(cfg: config.TTSConfig, text: str) -> bytes:
    key = _require_env("ELEVENLABS_API_KEY")
    voice = cfg.voice or "21m00Tcm4TlvDq8ikWAM"
    url = (f"https://api.elevenlabs.io/v1/text-to-speech/{voice}"
           f"?output_format=mp3_44100_128")
    body = {"text": text, "model_id": cfg.model or "eleven_multilingual_v2"}
    # ElevenLabs 返回原始 mp3 字节。
    return _http_post(url, body, {"xi-api-key": key, "Accept": "audio/mpeg"})


def _synth_fishaudio(cfg: config.TTSConfig, text: str) -> bytes:
    key = _require_env("FISH_API_KEY")
    body = {"text": text, "format": "mp3"}
    if cfg.voice:
        body["reference_id"] = cfg.voice
    # Fish Audio /v1/tts 接受 JSON,直接返回音频字节。
    return _http_post("https://api.fish.audio/v1/tts", body,
                      {"Authorization": f"Bearer {key}"})


def _synth_minimax(cfg: config.TTSConfig, text: str) -> bytes:
    key = _require_env("MINIMAX_API_KEY")
    group = _require_env("MINIMAX_GROUP_ID")
    url = f"https://api.minimax.chat/v1/t2a_v2?GroupId={group}"
    body = {
        "model": cfg.model or "speech-01-turbo",
        "text": text,
        "stream": False,
        "voice_setting": {"voice_id": cfg.voice, "speed": cfg.speed},
        "audio_setting": {"format": "mp3", "sample_rate": 32000},
    }
    raw = _http_post(url, body, {"Authorization": f"Bearer {key}"})
    data = json.loads(raw)
    # 返回 JSON,音频为 data.audio(hex 编码)。
    hexstr = (data.get("data") or {}).get("audio")
    if not hexstr:
        err = data.get("base_resp", {})
        raise RuntimeError(f"Minimax 无音频返回:{err or data}")
    return bytes.fromhex(hexstr)


def _synth_doubao(cfg: config.TTSConfig, text: str) -> bytes:
    import uuid
    appid = _require_env("DOUBAO_APP_ID")
    token = _require_env("DOUBAO_ACCESS_TOKEN")
    body = {
        "app": {"appid": appid, "token": token,
                "cluster": cfg.model or "volcano_tts"},
        "user": {"uid": "tts-quality-eval"},
        "audio": {"voice_type": cfg.voice, "encoding": "mp3",
                  "speed_ratio": cfg.speed},
        "request": {"reqid": str(uuid.uuid4()), "text": text, "operation": "query"},
    }
    # 火山引擎鉴权头是特殊的 'Bearer;{token}' 形式;音频为 base64 编码的 data 字段。
    raw = _http_post("https://openspeech.bytedance.com/api/v1/tts", body,
                     {"Authorization": f"Bearer;{token}"})
    data = json.loads(raw)
    b64 = data.get("data")
    if not b64:
        raise RuntimeError(f"豆包无音频返回:code={data.get('code')} "
                           f"message={data.get('message')}")
    return base64.b64decode(b64)


_SYNTH_DISPATCH = {
    "openai": _synth_openai,
    "elevenlabs": _synth_elevenlabs,
    "fishaudio": _synth_fishaudio,
    "minimax": _synth_minimax,
    "doubao": _synth_doubao,
}


# ---------------------------------------------------------------------------
# 2) 时长探测(ffprobe)
# ---------------------------------------------------------------------------
def probe_duration(path: str) -> float:
    """返回音频时长(秒)。ffprobe 缺失或出错时抛异常。"""
    if shutil.which("ffprobe") is None:
        raise RuntimeError("未找到 ffprobe,请安装 ffmpeg(macOS: brew install ffmpeg)。")
    proc = subprocess.run(
        ["ffprobe", "-v", "error", "-show_entries", "format=duration",
         "-of", "default=noprint_wrappers=1:nokey=1", path],
        capture_output=True, text=True,
    )
    if proc.returncode != 0:
        raise RuntimeError(f"ffprobe 失败: {proc.stderr.strip()}")
    out = proc.stdout.strip()
    try:
        return float(out)
    except ValueError:
        raise RuntimeError(f"ffprobe 输出无法解析为时长: {out!r}")


# ---------------------------------------------------------------------------
# 3) 回译(Whisper 转写)
# ---------------------------------------------------------------------------
# 用简体中文提示语引导 Whisper 输出简体,避免它偶尔返回繁体导致 CER 被字形差异
# 虚高(那是转写脚本选择问题,不是 TTS 发音错误)。
_ZH_PROMPT = "以下是普通话简体中文的句子。"


def transcribe(path: str) -> str:
    with open(path, "rb") as f:
        tr = get_client().audio.transcriptions.create(
            model=config.WHISPER_MODEL, file=f, language="zh", prompt=_ZH_PROMPT,
        )
    return tr.text or ""


# ---------------------------------------------------------------------------
# 4) 文本归一化 + 字错误率(中文用字级 CER,等价于书中所说 WER 的可懂度维度)
# ---------------------------------------------------------------------------
def normalize(text: str) -> str:
    """去掉标点/空白,只保留 CJK / 字母 / 数字,并小写,便于逐字比较。"""
    text = text.lower()
    return "".join(ch for ch in text if ch.isalnum())


def _edit_distance(a: str, b: str) -> int:
    """Levenshtein 距离(字符级)。"""
    if a == b:
        return 0
    if not a:
        return len(b)
    if not b:
        return len(a)
    prev = list(range(len(b) + 1))
    for i, ca in enumerate(a, 1):
        cur = [i]
        for j, cb in enumerate(b, 1):
            cur.append(min(
                prev[j] + 1,        # 删除
                cur[j - 1] + 1,     # 插入
                prev[j - 1] + (ca != cb),  # 替换
            ))
        prev = cur
    return prev[-1]


@dataclass
class ErrorRate:
    cer: float          # 字错误率 = 编辑距离 / 参考字数
    accuracy: float     # 字准确率 = 1 - cer(下限 0)
    edits: int
    ref_len: int


def char_error_rate(reference: str, hypothesis: str) -> ErrorRate:
    ref = normalize(reference)
    hyp = normalize(hypothesis)
    if not ref:
        return ErrorRate(0.0, 1.0, 0, 0)
    dist = _edit_distance(ref, hyp)
    cer = dist / len(ref)
    return ErrorRate(cer=cer, accuracy=max(0.0, 1.0 - cer), edits=dist, ref_len=len(ref))


# ---------------------------------------------------------------------------
# 5) LLM Rubric 评审(默认,OpenAI 闭环)
# ---------------------------------------------------------------------------
RUBRIC_DIMENSIONS = ["清晰度", "自然度", "停顿节奏", "整体"]

# 维度说明(供 --dump-rubric 离线打印,也是评审 prompt 的依据)。括号内标注与书中
# 四维度(准确性 / 自然度 / 情感表达 / 音色一致性)的对应关系。
RUBRIC_DESCRIPTIONS = {
    "清晰度": "转写与原文是否高度一致,漏字/错字/多字越多分越低(对应书中「准确性」)。",
    "自然度": "语速是否接近自然朗读(中文约 4-6 字/秒),过快>7 或过慢<3 都不自然。",
    "停顿节奏": "结合语速与文本长度判断停顿/节奏是否合理,过快通常意味吞字、节奏差。",
    "整体": "综合以上给出的总体印象分。",
}
# 说明:默认(回译)评审看不到音频,无法覆盖书中「情感表达 / 音色一致性」;这两维需
# 多模态直接听音频,用 --gemini 复现(音色一致性还需参考语音,本 demo 未提供)。

_JUDGE_SYSTEM = """你是严格的 TTS(文本转语音)质量评审专家。
你将拿到:原始参考文本、该文本的期望情感、由 Whisper 对合成语音回译得到的转写文本,
以及从音频客观测得的时长、语速(字/秒)和字错误率(CER)。
请据此对合成语音质量按 Rubric 逐维度打分(1-5 的整数,5 最好):

- 清晰度:转写与原文是否高度一致(漏字/错字/多字越多分越低;CER 越高分越低)。
- 自然度:语速是否接近自然朗读(中文自然朗读约 4-6 字/秒;过快>7 或过慢<3 都不自然)。
- 停顿节奏:结合语速与文本长度,判断停顿/节奏是否合理(过快通常意味着吞字、节奏差)。
- 整体:综合以上给出的总体印象分。

注意:你看不到音频本身,只能基于以上可测特征做保守、可解释的判断。
只输出 JSON,格式:
{"清晰度": {"score": int, "reason": str},
 "自然度": {"score": int, "reason": str},
 "停顿节奏": {"score": int, "reason": str},
 "整体": {"score": int, "reason": str}}
reason 用一句简短中文说明。"""


@dataclass
class RubricResult:
    scores: dict            # 维度 -> int
    reasons: dict           # 维度 -> str
    raw: str = ""


def judge_rubric(reference: str, emotion: str, hypothesis: str,
                 duration: float, cer: float, model: Optional[str] = None) -> RubricResult:
    """用评审模型(默认 gpt-5.6-luna)按 Rubric 打分。返回结构化分数 + 点评。

    评审 chat 调用支持 OpenRouter 回退(见 get_judge_client_and_model)。"""
    chars = len(normalize(reference))
    speed = chars / duration if duration > 0 else 0.0
    user = (
        f"原始参考文本:{reference}\n"
        f"期望情感:{emotion}\n"
        f"Whisper 回译文本:{hypothesis}\n"
        f"音频时长:{duration:.2f}\n"
        f"语速:{speed:.2f} 字/秒(参考文本 {chars} 个有效字符)\n"
        f"字错误率 CER:{cer:.3f}\n"
    )
    judge_client, judge_model = get_judge_client_and_model(model or config.JUDGE_MODEL)
    resp = judge_client.chat.completions.create(
        model=judge_model,
        messages=[{"role": "system", "content": _JUDGE_SYSTEM},
                  {"role": "user", "content": user}],
        temperature=0.0,
        response_format={"type": "json_object"},
    )
    raw = resp.choices[0].message.content or "{}"
    data = json.loads(raw)
    scores, reasons = {}, {}
    for dim in RUBRIC_DIMENSIONS:
        item = data.get(dim, {})
        if isinstance(item, dict):
            scores[dim] = int(item.get("score") or 0)   # score 缺失或为 null 时按 0 分
            reasons[dim] = str(item.get("reason", "")).strip()
        else:  # 兼容模型直接返回数字(null 按 0 分)
            scores[dim] = int(item or 0)
            reasons[dim] = ""
    return RubricResult(scores=scores, reasons=reasons, raw=raw)


# ---------------------------------------------------------------------------
# 6) 可选:Gemini 多模态音频评审(书中方案)。用 REST,避免额外 SDK 依赖。
# ---------------------------------------------------------------------------
def _resolve_gemini_model(api_key: str) -> str:
    """探测当前可用的 Gemini 模型,避免默认名过期。"""
    import urllib.request
    url = f"https://generativelanguage.googleapis.com/v1beta/models?key={api_key}"
    try:
        with urllib.request.urlopen(url, timeout=20) as r:
            data = json.loads(r.read())
        names = [m["name"].split("/")[-1] for m in data.get("models", [])
                 if "generateContent" in m.get("supportedGenerationMethods", [])]
        # 优先默认的 gemini-3.5-flash(已验证支持音频输入),再退到 pro / 旧 flash 系列。
        for want in (config.GEMINI_MODEL_DEFAULT, "gemini-3.5-flash",
                     "gemini-2.5-pro", "gemini-2.5-flash", "gemini-flash-latest"):
            if want in names:
                return want
        # 退而求其次:任意非 tts/image 的可用模型
        for n in names:
            if "tts" not in n and "image" not in n and "embedding" not in n:
                return n
    except Exception:
        pass
    return config.GEMINI_MODEL_DEFAULT


def judge_gemini_audio(reference: str, emotion: str, audio_path: str) -> RubricResult:
    """把合成音频 + 原文 + Rubric 一起交给 Gemini 多模态直接「听」并打分。

    需要 GEMINI_API_KEY。默认关闭;--gemini 开启。失败抛异常由上层记为失败。
    """
    import urllib.request
    key = os.environ.get("GEMINI_API_KEY", "").strip()
    if not key:
        raise RuntimeError("缺少 GEMINI_API_KEY,无法使用 Gemini 音频评审。")
    model = _resolve_gemini_model(key)
    with open(audio_path, "rb") as f:
        audio_b64 = base64.b64encode(f.read()).decode()
    prompt = (
        "你是 TTS 质量评审专家。请直接聆听下面这段合成语音,对照原始文本与期望情感,"
        "按 1-5 分为四个维度打分并给出简短理由,只输出 JSON:"
        '{"清晰度":{"score":int,"reason":str},"自然度":{"score":int,"reason":str},'
        '"停顿节奏":{"score":int,"reason":str},"整体":{"score":int,"reason":str}}\n'
        f"原始文本:{reference}\n期望情感:{emotion}"
    )
    body = {
        "contents": [{"parts": [
            {"text": prompt},
            {"inline_data": {"mime_type": "audio/mp3", "data": audio_b64}},
        ]}],
        "generationConfig": {"temperature": 0.0, "responseMimeType": "application/json"},
    }
    url = (f"https://generativelanguage.googleapis.com/v1beta/models/"
           f"{model}:generateContent?key={key}")
    req = urllib.request.Request(
        url, data=json.dumps(body).encode(),
        headers={"Content-Type": "application/json"}, method="POST",
    )
    with urllib.request.urlopen(req, timeout=90) as r:
        data = json.loads(r.read())
    # Gemini 在安全拦截时不返回 candidates(或 candidate 无 content/parts),
    # 防御式取值并给出带 promptFeedback 的清晰错误,交由上层记为该条失败。
    candidates = data.get("candidates") or []
    parts = []
    if candidates:
        parts = (candidates[0].get("content") or {}).get("parts") or []
    if not parts or not parts[0].get("text"):
        raise RuntimeError(f"Gemini 未返回评审文本:{data.get('promptFeedback') or data}")
    text = parts[0]["text"]
    parsed = json.loads(text)
    scores, reasons = {}, {}
    # 评审 JSON 的 score 字段缺失或为 null 时按 0 分处理,与 judge_rubric 一致
    for dim in RUBRIC_DIMENSIONS:
        item = parsed.get(dim, {})
        scores[dim] = int(item.get("score") or 0) if isinstance(item, dict) else int(item or 0)
        reasons[dim] = str(item.get("reason", "")).strip() if isinstance(item, dict) else ""
    return RubricResult(scores=scores, reasons=reasons, raw=text)

test_judge_robustness.py

"""
Regression tests for judge-response robustness (实验 6-5 TTS 质量评估).

Covers two failure classes on LLM/Gemini judge responses:
  - judge_rubric: judge returns "score": null (or a bare null dimension) -> int(None) TypeError
  - judge_gemini_audio: safety-blocked Gemini responses have no
    candidates/content/parts -> KeyError/IndexError instead of a clear error

Network is stubbed: the OpenAI-compatible judge client is replaced with a fake,
and urllib.request.urlopen is monkeypatched for the Gemini REST call.
"""
import io
import json

import pytest

import pipeline


class _FakeMessage:
    content = "{}"


class _FakeChoice:
    message = _FakeMessage()


class _FakeResp:
    choices = [_FakeChoice()]


class _FakeCompletions:
    @staticmethod
    def create(**kwargs):
        return _FakeResp()


class _FakeChat:
    completions = _FakeCompletions()


class _FakeClient:
    chat = _FakeChat()


def _stub_judge(monkeypatch, payload: dict):
    _FakeMessage.content = json.dumps(payload, ensure_ascii=False)
    monkeypatch.setattr(
        pipeline, "get_judge_client_and_model", lambda model=None: (_FakeClient(), "fake-judge"))


def test_judge_rubric_tolerates_null_score(monkeypatch):
    """'score': null in a dimension dict is scored 0, not int(None) TypeError."""
    _stub_judge(monkeypatch, {
        "清晰度": {"score": None, "reason": "无法判断"},
        "自然度": {"score": 4, "reason": "语速正常"},
        "停顿节奏": {"score": 3},
        "整体": {"score": 5, "reason": "总体可用"},
    })
    rub = pipeline.judge_rubric("原文文本", "中性", "回译文本", 3.0, 0.05)
    assert rub.scores["清晰度"] == 0
    assert rub.scores["自然度"] == 4
    assert rub.scores["整体"] == 5


def test_judge_rubric_tolerates_null_dimension(monkeypatch):
    """A bare null dimension (non-dict) is scored 0, not int(None) TypeError."""
    _stub_judge(monkeypatch, {"清晰度": None, "自然度": 4, "停顿节奏": 3, "整体": 5})
    rub = pipeline.judge_rubric("原文文本", "中性", "回译文本", 3.0, 0.05)
    assert rub.scores["清晰度"] == 0
    assert rub.scores["自然度"] == 4


class _FakeHTTPResp(io.BytesIO):
    def __enter__(self):
        return self

    def __exit__(self, *args):
        return False


def _stub_gemini(monkeypatch, payload: dict):
    monkeypatch.setenv("GEMINI_API_KEY", "fake-key-for-test")
    monkeypatch.setattr(pipeline, "_resolve_gemini_model", lambda key: "gemini-fake")
    monkeypatch.setattr("urllib.request.urlopen",
                        lambda req, timeout=None: _FakeHTTPResp(json.dumps(payload).encode()))


@pytest.mark.parametrize("payload", [
    {"promptFeedback": {"blockReason": "SAFETY"}},      # prompt 被拦截:无 candidates
    {"candidates": []},                                  # 生成被拦截:空 candidates
    {"candidates": [{"finishReason": "SAFETY", "index": 0}]},  # candidate 无 content
])
def test_judge_gemini_audio_blocked_raises_clear_error(monkeypatch, tmp_path, payload):
    """Blocked/empty Gemini responses raise a clear RuntimeError, not KeyError/IndexError."""
    _stub_gemini(monkeypatch, payload)
    audio = tmp_path / "a.mp3"
    audio.write_bytes(b"\xff\xfb" + b"\x00" * 256)
    with pytest.raises(RuntimeError, match="Gemini 未返回评审文本"):
        pipeline.judge_gemini_audio("原文", "中性", str(audio))


def test_judge_gemini_audio_parses_valid_response(monkeypatch, tmp_path):
    """A normal Gemini response still parses (defensive navigation keeps working)."""
    inner = json.dumps({"清晰度": {"score": 4, "reason": "ok"}, "自然度": 4,
                        "停顿节奏": None, "整体": {"score": 5}}, ensure_ascii=False)
    _stub_gemini(monkeypatch, {
        "candidates": [{"content": {"parts": [{"text": inner}]}}],
    })
    audio = tmp_path / "a.mp3"
    audio.write_bytes(b"\xff\xfb" + b"\x00" * 256)
    rub = pipeline.judge_gemini_audio("原文", "中性", str(audio))
    assert rub.scores["清晰度"] == 4
    assert rub.scores["停顿节奏"] == 0   # null score -> 0
    assert rub.scores["整体"] == 5


if __name__ == "__main__":
    pytest.main([__file__, "-v"])