test: 按模块重写测试代码,删除旧平铺结构
按"测试规则"重写 tests/:一个模块一个目录、用例按数据→过程→验证三段书写、 不保留全局 conftest.py、测试过程只调用真实生产代码。 结构(73 个文件、30 个模块目录、477 用例): - tests/nodes/ 15 个模块目录(srt/whisper/ass/ffmpeg/frame_extract/vlm/ subtitle_ocr/llm/llm_filter/subtitle_cleanup/subtitle_correction/ proper_nouns/adaptive_pool/vad_profiler/echo); - tests/app/ 11 个模块目录(db/scheduler/batch/maintenance/registry/seed/ storage/config/logging/main/routers 三组 API); - tests/sdk/test_models、tests/web/test_crop、tests/shared(公共设施)。 测试数据随模块目录入库(tests/**/data/),删除根级 testdata/;.gitignore 的 data/ 改为 /data/,否则会连带忽略 tests/**/data/ 导致测试数据无法入库。 顺带发现并修复三个真实缺陷: - nodes/srt.py:相邻条目缺少空行时把下一条时间轴吞进正文(静默错位), 改为正文行遇时间戳行即报错; - src/wov_app/scheduler.py:_file_size 只捕获 OSError,含 \x00 的产物 URI 抛 ValueError 导致任务误判失败,改为同时捕获; - nodes/subtitle_correction.py:生产代码依赖测试包解析 SRT, 改用生产模块 nodes/srt.py。 真实模型/服务集成测试按外部状态跳过:新增 tests/shared/gpu_memory.py (运行时探测显存、CUDA OOM 转跳过)与 tests/shared/llm_service.py (无 Key / 余额 / 限流转跳过)。全量 477 passed。
This commit is contained in:
@@ -0,0 +1,416 @@
|
||||
"""nodes/subtitle_correction.py 的模块级测试(数据 → 测试过程 → 验证结果)。
|
||||
|
||||
被测模块:`nodes/subtitle_correction.py`(字幕领域纠错:上下文过滤 + LLM
|
||||
推断误听词),可独立调用。网络属于允许 mock 的 I/O 边界;集成用例调用真实
|
||||
LLM 验证误听泛化能力。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from nodes.subtitle_correction import (
|
||||
CONTEXT_WINDOW,
|
||||
PROPER_SESSION_WORDS,
|
||||
_build_context,
|
||||
_extract_target_line,
|
||||
_is_fragment,
|
||||
_read_srt_entries,
|
||||
_serialize_srt,
|
||||
_system_prompt,
|
||||
correct_entry,
|
||||
invoke,
|
||||
)
|
||||
from wov_sdk.models import InvokeRequest
|
||||
|
||||
# 模块专用数据目录(真实素材缺失时相关用例跳过)。
|
||||
DATA_DIR = Path(__file__).resolve().parent / "data"
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
"""假 HTTP 响应:返回给定的 content 字符串。"""
|
||||
|
||||
def __init__(self, content: str) -> None:
|
||||
self._payload = {
|
||||
"choices": [{"message": {"content": content}}],
|
||||
}
|
||||
|
||||
def read(self) -> bytes:
|
||||
return json.dumps(self._payload, ensure_ascii=False).encode("utf-8")
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _srt(*cues: tuple[float, float, str]) -> str:
|
||||
"""把 (起始秒, 结束秒, 文本) 拼成标准 SRT 文本。"""
|
||||
|
||||
def fmt(seconds: float) -> str:
|
||||
hours, rest = divmod(seconds, 3600)
|
||||
minutes, secs = divmod(rest, 60)
|
||||
return f"{int(hours):02d}:{int(minutes):02d}:{int(secs):02d},{int(round((secs - int(secs)) * 1000)):03d}"
|
||||
|
||||
return "\n".join(
|
||||
f"{i}\n{fmt(start)} --> {fmt(end)}\n{text}\n"
|
||||
for i, (start, end, text) in enumerate(cues, 1)
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 碎片识别与上下文构造
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_is_fragment_detects_pure_phrases() -> None:
|
||||
"""纯语气词/单字为碎片;有实义的句子不是。"""
|
||||
# 数据:碎片与正常句子。
|
||||
# 测试过程与验证结果
|
||||
assert _is_fragment("あ") is True
|
||||
assert _is_fragment("ん?") is True
|
||||
assert _is_fragment("はい") is True
|
||||
assert _is_fragment("") is True
|
||||
assert _is_fragment("それじゃあ、始めましょう") is False
|
||||
|
||||
|
||||
def test_build_context_filters_fragments_but_keeps_target() -> None:
|
||||
"""上下文过滤语气词碎片,但目标条目始终保留并标记。"""
|
||||
# 数据:目标条目周围有碎片与正常句。
|
||||
entries = [
|
||||
{"start": 100.0, "end": 101.0, "text": "あ"},
|
||||
{"start": 101.0, "end": 104.0, "text": "それじゃあ、始めましょう"},
|
||||
{"start": 104.0, "end": 107.0, "text": "相手がバナナをしゃぶってくれて"},
|
||||
{"start": 107.0, "end": 108.0, "text": "ん"},
|
||||
]
|
||||
|
||||
# 测试过程:以索引 2 为目标。
|
||||
context = _build_context(entries, 2)
|
||||
|
||||
# 验证结果:碎片不出现,目标带标记,正常上下文保留。
|
||||
assert "相手がバナナをしゃぶってくれて <-- 目标" in context
|
||||
assert "それじゃあ" in context
|
||||
assert "あ\n" not in context and "[100.00] あ" not in context
|
||||
|
||||
|
||||
def test_build_context_respects_time_window() -> None:
|
||||
"""窗口外的条目不进上下文(默认 ±60 秒)。"""
|
||||
# 数据:目标与远处条目。
|
||||
entries = [
|
||||
{"start": 0.0, "end": 2.0, "text": "很早之前说的话"},
|
||||
{"start": 500.0, "end": 502.0, "text": "目标所在位置"},
|
||||
{"start": 1000.0, "end": 1002.0, "text": "很久之后说的话"},
|
||||
]
|
||||
|
||||
# 测试过程
|
||||
context = _build_context(entries, 1)
|
||||
|
||||
# 验证结果:只有目标出现。
|
||||
assert "目标所在位置" in context
|
||||
assert "很早之前" not in context
|
||||
assert "很久之后" not in context
|
||||
assert CONTEXT_WINDOW == 60
|
||||
|
||||
|
||||
def test_serialize_srt_round_trip() -> None:
|
||||
"""序列化输出合法 SRT(时间戳格式正确、条数一致)。"""
|
||||
# 数据:两个条目。
|
||||
entries = [
|
||||
{"start": 1.0, "end": 2.5, "text": "第一句"},
|
||||
{"start": 3.0, "end": 4.0, "text": "第二句"},
|
||||
]
|
||||
|
||||
# 测试过程
|
||||
text = _serialize_srt(entries)
|
||||
|
||||
# 验证结果:序号、时间戳格式与正文。
|
||||
assert text.startswith("1\n00:00:01,000 --> 00:00:02,500\n第一句")
|
||||
assert "\n2\n" in text
|
||||
assert text.count("-->") == 2
|
||||
|
||||
|
||||
def test_read_srt_entries_uses_shared_parser(tmp_path: Path) -> None:
|
||||
"""读取真实 SRT 文件返回带秒级时间轴的条目。"""
|
||||
# 数据:真实文件。
|
||||
srt_path = tmp_path / "in.srt"
|
||||
srt_path.write_text(_srt((1.0, 2.0, "你好")), encoding="utf-8")
|
||||
|
||||
# 测试过程
|
||||
entries = _read_srt_entries(srt_path)
|
||||
|
||||
# 验证结果
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["text"] == "你好"
|
||||
assert entries[0]["start"] == 1.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 提示词
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_system_prompt_contains_domain_terms_and_no_mishearing_examples() -> None:
|
||||
"""系统提示词含领域词表但不含具体误听例子(避免过拟合)。"""
|
||||
# 数据:目标语言 zh-CN。
|
||||
# 测试过程
|
||||
prompt = _system_prompt("zh-CN")
|
||||
|
||||
# 验证结果:包含领域词(如 チンポ),且要求按上下文推断而非照搬字面。
|
||||
assert "zh-CN" in prompt
|
||||
assert any(term in prompt for term in PROPER_SESSION_WORDS)
|
||||
assert "不要机械照搬字面词" in prompt
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# correct_entry / _extract_target_line
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_extract_target_line_matches_by_time() -> None:
|
||||
"""按时间戳匹配目标行(容忍 LLM 输出的编号差异)。"""
|
||||
# 数据:LLM 输出两行,目标时间 120.50。
|
||||
content = "120.00 这是上一句\n120.50 这是目标句\n"
|
||||
|
||||
# 测试过程
|
||||
target = _extract_target_line(content, 120.5)
|
||||
|
||||
# 验证结果
|
||||
assert target == "这是目标句"
|
||||
|
||||
|
||||
def test_extract_target_line_picks_closest_timestamp() -> None:
|
||||
"""多行输出时取时间戳最接近目标的那一行(不依赖行顺序)。"""
|
||||
# 数据:三行,中间一行最接近 120.50。
|
||||
content = "100.00 甲\n120.40 目标句\n200.00 乙\n"
|
||||
|
||||
# 测试过程
|
||||
target = _extract_target_line(content, 120.5)
|
||||
|
||||
# 验证结果
|
||||
assert target == "目标句"
|
||||
|
||||
|
||||
def test_extract_target_line_returns_empty_without_timestamps() -> None:
|
||||
"""输出行不带时间戳前缀时无法定位目标,返回空串(不猜)。"""
|
||||
# 数据:单行纯译文,无时间戳。
|
||||
content = "目标译文\n"
|
||||
|
||||
# 测试过程与验证结果
|
||||
assert _extract_target_line(content, 120.5) == ""
|
||||
|
||||
|
||||
def test_extract_target_line_ignores_unparseable_lines() -> None:
|
||||
"""无法解析时间戳的行被跳过,仍能取到最接近的可解析行。"""
|
||||
# 数据:首行无时间戳,次行有。
|
||||
content = "这是解释性文字\n120.50 真正的译文\n"
|
||||
|
||||
# 测试过程与验证结果
|
||||
assert _extract_target_line(content, 120.5) == "真正的译文"
|
||||
|
||||
|
||||
def test_correct_entry_sends_context_and_returns_translation(monkeypatch) -> None:
|
||||
"""correct_entry 把目标前后上下文发给 LLM,返回目标条目译文。"""
|
||||
# 数据:真实形态的条目列表(含误听词)。
|
||||
entries = [
|
||||
{"start": 100.0, "end": 103.0, "text": "気持ちいいところに当たってるね"},
|
||||
{"start": 103.0, "end": 106.0, "text": "そろそろマンゴーが濡れてきました"},
|
||||
]
|
||||
captured: dict = {}
|
||||
|
||||
def fake_urlopen(http_request, timeout=None):
|
||||
captured["body"] = json.loads(http_request.data.decode("utf-8"))
|
||||
return _FakeResponse("103.00 那里已经湿了呢\n")
|
||||
|
||||
monkeypatch.setenv("LLM_API_KEY", "sk-test")
|
||||
monkeypatch.setattr(urllib.request, "urlopen", fake_urlopen)
|
||||
|
||||
# 测试过程
|
||||
result = correct_entry(entries[1], entries, 1, {"target_language": "zh-CN"})
|
||||
|
||||
# 验证结果:返回译文,且请求体带上下文与系统提示词。
|
||||
assert result == "那里已经湿了呢"
|
||||
body = captured["body"]
|
||||
assert body["messages"][0]["role"] == "system"
|
||||
assert "マンゴー" in body["messages"][1]["content"]
|
||||
|
||||
|
||||
def test_correct_entry_uses_default_model_when_not_configured(monkeypatch) -> None:
|
||||
"""未指定模型/环境变量时使用节点自有兜底模型(与全局 LLM_MODEL 解耦)。"""
|
||||
# 数据:清空环境变量,捕获请求体。
|
||||
entries = [{"start": 1.0, "end": 2.0, "text": "テスト"}]
|
||||
captured: dict = {}
|
||||
monkeypatch.delenv("LLM_MODEL", raising=False)
|
||||
monkeypatch.setenv("LLM_API_KEY", "sk-test")
|
||||
monkeypatch.setattr(
|
||||
urllib.request, "urlopen",
|
||||
lambda req, timeout=None: (captured.update(json.loads(req.data.decode("utf-8"))),
|
||||
_FakeResponse("1.00 测试"))[1],
|
||||
)
|
||||
|
||||
# 测试过程
|
||||
correct_entry(entries[0], entries, 0, {})
|
||||
|
||||
# 验证结果:模型名固定为节点兜底(Qwen3.6 系列,有意不跟随全局默认)。
|
||||
assert captured["model"] == "Qwen/Qwen3.6-35B-A3B"
|
||||
|
||||
|
||||
def test_correct_entry_param_model_wins(monkeypatch) -> None:
|
||||
"""参数 model 优先于兜底值。"""
|
||||
# 数据:传入自定义模型。
|
||||
entries = [{"start": 1.0, "end": 2.0, "text": "テスト"}]
|
||||
captured: dict = {}
|
||||
monkeypatch.setenv("LLM_API_KEY", "sk-test")
|
||||
monkeypatch.setattr(
|
||||
urllib.request, "urlopen",
|
||||
lambda req, timeout=None: (captured.update(json.loads(req.data.decode("utf-8"))),
|
||||
_FakeResponse("1.00 测试"))[1],
|
||||
)
|
||||
|
||||
# 测试过程
|
||||
correct_entry(entries[0], entries, 0, {"model": "自定义/纠错模型"})
|
||||
|
||||
# 验证结果
|
||||
assert captured["model"] == "自定义/纠错模型"
|
||||
|
||||
|
||||
def test_correct_entry_returns_empty_on_network_error(monkeypatch) -> None:
|
||||
"""网络错误时返回空字符串(调用方按"未纠错"处理,不打断整节点)。"""
|
||||
# 数据:urlopen 抛错。
|
||||
entries = [{"start": 1.0, "end": 2.0, "text": "テスト"}]
|
||||
monkeypatch.setenv("LLM_API_KEY", "sk-test")
|
||||
|
||||
def fake_urlopen(*a, **k):
|
||||
raise urllib.error.URLError("boom")
|
||||
|
||||
monkeypatch.setattr(urllib.request, "urlopen", fake_urlopen)
|
||||
|
||||
# 测试过程与验证结果
|
||||
assert correct_entry(entries[0], entries, 0, {}) == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# invoke 全流程
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_invoke_writes_corrected_srt(monkeypatch, tmp_path: Path) -> None:
|
||||
"""invoke 读取 SRT、逐条纠错并写出 corrected.srt。"""
|
||||
# 数据:两条字幕,LLM 每次都返回目标译文。
|
||||
srt_path = tmp_path / "asr.srt"
|
||||
srt_path.write_text(_srt((1.0, 2.0, "こんにちは"), (3.0, 4.0, "さようなら")), encoding="utf-8")
|
||||
monkeypatch.setenv("LLM_API_KEY", "sk-test")
|
||||
|
||||
def fake_urlopen(http_request, timeout=None):
|
||||
body = json.loads(http_request.data.decode("utf-8"))
|
||||
# 回显目标时间戳(模拟真实模型按约定格式输出)。
|
||||
user = body["messages"][1]["content"]
|
||||
stamp = user.split("[")[1].split("]")[0].strip()
|
||||
return _FakeResponse(f"{stamp} 纠错后的译文\n")
|
||||
|
||||
monkeypatch.setattr(urllib.request, "urlopen", fake_urlopen)
|
||||
request = InvokeRequest(
|
||||
run_id="r", node_instance_id="n", params={},
|
||||
inputs={"srt_uri": str(srt_path)}, output_dir=str(tmp_path / "out"),
|
||||
)
|
||||
|
||||
# 测试过程
|
||||
response = invoke(request)
|
||||
|
||||
# 验证结果:产物存在,时间轴保留,正文被替换。
|
||||
assert response.status == "completed", response.error
|
||||
content = Path(response.outputs["srt_uri"]).read_text(encoding="utf-8")
|
||||
assert "00:00:01,000 --> 00:00:02,000" in content
|
||||
assert "纠错后的译文" in content
|
||||
assert content.count("-->") == 2
|
||||
|
||||
|
||||
def test_invoke_keeps_original_when_correction_empty(monkeypatch, tmp_path: Path) -> None:
|
||||
"""纠错返回空时保留原始文本(不产出空字幕)。"""
|
||||
# 数据:LLM 返回空内容。
|
||||
srt_path = tmp_path / "asr.srt"
|
||||
srt_path.write_text(_srt((1.0, 2.0, "原文内容")), encoding="utf-8")
|
||||
monkeypatch.setenv("LLM_API_KEY", "sk-test")
|
||||
monkeypatch.setattr(
|
||||
urllib.request, "urlopen",
|
||||
lambda *a, **k: _FakeResponse(""),
|
||||
)
|
||||
request = InvokeRequest(
|
||||
run_id="r", node_instance_id="n", params={},
|
||||
inputs={"srt_uri": str(srt_path)}, output_dir=str(tmp_path / "out"),
|
||||
)
|
||||
|
||||
# 测试过程
|
||||
response = invoke(request)
|
||||
|
||||
# 验证结果
|
||||
assert response.status == "completed"
|
||||
content = Path(response.outputs["srt_uri"]).read_text(encoding="utf-8")
|
||||
assert "原文内容" in content
|
||||
|
||||
|
||||
def test_invoke_fails_without_input(tmp_path: Path) -> None:
|
||||
"""缺少 srt_uri 时失败。"""
|
||||
# 数据:空输入。
|
||||
request = InvokeRequest(
|
||||
run_id="r", node_instance_id="n", params={}, inputs={}, output_dir=str(tmp_path)
|
||||
)
|
||||
|
||||
# 测试过程
|
||||
response = invoke(request)
|
||||
|
||||
# 验证结果
|
||||
assert response.status == "failed"
|
||||
assert "srt_uri" in (response.error or "")
|
||||
|
||||
|
||||
def test_invoke_fails_when_input_missing(tmp_path: Path) -> None:
|
||||
"""输入文件不存在时失败。"""
|
||||
# 数据:不存在的路径。
|
||||
request = InvokeRequest(
|
||||
run_id="r", node_instance_id="n", params={},
|
||||
inputs={"srt_uri": str(tmp_path / "nope.srt")}, output_dir=str(tmp_path),
|
||||
)
|
||||
|
||||
# 测试过程
|
||||
response = invoke(request)
|
||||
|
||||
# 验证结果
|
||||
assert response.status == "failed"
|
||||
assert "not found" in (response.error or "")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 真实 LLM 集成:误听泛化
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_real_llm_generalizes_to_unseen_mishearing() -> None:
|
||||
"""真实 LLM 校准:未在提示词中出现的误听词也能结合上下文正确推断。"""
|
||||
# 数据:模拟 ASR 把 チンポ/マンコ 听成 バナナ/マンゴー。
|
||||
from tests.shared.llm_service import probe_llm_or_skip
|
||||
|
||||
# 先探针真实服务:不可用(无 Key / 余额 / 限流)时跳过,避免把外部
|
||||
# 状态问题误判成"模型未泛化"(correct_entry 会把调用异常吞成空串)。
|
||||
probe_llm_or_skip()
|
||||
entries = [
|
||||
{"start": 1200.0, "end": 1203.0, "text": "相手がバナナをしゃぶってくれて"},
|
||||
{"start": 1203.0, "end": 1206.0, "text": "そろそろマンゴーが濡れてきました"},
|
||||
{"start": 1206.0, "end": 1209.0, "text": "気持ちいいところに当たってるね"},
|
||||
{"start": 1220.0, "end": 1224.0, "text": "もっとマンゴーを舐めてください"},
|
||||
{"start": 1224.0, "end": 1226.0, "text": "いっぱい出してね"},
|
||||
]
|
||||
|
||||
# 测试过程:对含误听词的第二条纠错。
|
||||
target = correct_entry(entries[1], entries, 1, {"target_language": "zh-CN"})
|
||||
|
||||
# 验证结果:输出体现性器官语义而非字面"芒果"。
|
||||
flagged = [k for k in ("肉棒", "鸡巴", "阴部", "小穴", "敏感", "那里", "湿") if k in target]
|
||||
assert flagged, f"泛化失败:模型仍字面直译,输出'{target}'"
|
||||
Reference in New Issue
Block a user