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:
2026-09-13 15:40:56 +08:00
parent 966f3e6b4b
commit 8a715a8064
139 changed files with 20810 additions and 10733 deletions
@@ -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}'"