Files
vrsub/tests/test_nodes.py
T

1565 lines
58 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""进程内节点测试。
合并原 5 个节点仓库的测试:echo、ffmpeg、whisper、llm、ass。节点已变为
进程内模块(nodes/*.py),移除了原入口点(__main__/run_node)相关测试。
"""
import base64
import json
import subprocess
import sys
import threading
import types
import urllib.error
import urllib.request
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from nodes.ass import invoke as ass_invoke
from nodes.ass import parse_srt, write_ass
from nodes.echo import invoke as echo_invoke
from nodes.ffmpeg import _ffmpeg_bin, invoke as ffmpeg_invoke
from nodes.llm import invoke as llm_invoke
from nodes.llm import translate_lines
from nodes.whisper import format_timestamp
from nodes.whisper import invoke as whisper_invoke
from nodes.whisper import resolve_model_path
from wov_sdk.models import InvokeRequest
# 单体根目录:tests/ 的上一级。
WORKSPACE = Path(__file__).resolve().parent.parent
# ---------------------------------------------------------------------------
# Echo 节点
# ---------------------------------------------------------------------------
def test_echo_invoke_text(tmp_path) -> None:
"""验证直接传入 text 时 echo 节点原样写出文本。"""
response = echo_invoke(
InvokeRequest(
run_id="run_1",
node_instance_id="ni_1",
inputs={"text": "hello echo"},
output_dir=str(tmp_path),
)
)
assert response.status == "completed"
assert response.outputs["text"] == "hello echo"
assert Path(response.outputs["file_uri"]).read_text(encoding="utf-8") == "hello echo"
def test_echo_invoke_absolute_file(tmp_path) -> None:
"""验证绝对路径 file_uri 的文件内容被读取为输入。"""
source = tmp_path / "input.txt"
source.write_text("from file", encoding="utf-8")
response = echo_invoke(
InvokeRequest(
run_id="run_2",
node_instance_id="ni_2",
inputs={"file_uri": str(source)},
output_dir=str(tmp_path / "out"),
)
)
assert response.status == "completed"
assert response.outputs["text"] == "from file"
def test_echo_invoke_relative_file(tmp_path) -> None:
"""验证相对路径 file_uri 以单体根目录为基准解析。"""
source = WORKSPACE / "relative_input.txt"
source.write_text("relative", encoding="utf-8")
try:
response = echo_invoke(
InvokeRequest(
run_id="run_3",
node_instance_id="ni_3",
inputs={"file_uri": "relative_input.txt"},
output_dir=str(tmp_path / "out"),
)
)
finally:
source.unlink()
assert response.status == "completed"
assert response.outputs["text"] == "relative"
def test_echo_invoke_default_text(tmp_path) -> None:
"""验证无任何输入时 echo 节点返回默认文本。"""
response = echo_invoke(
InvokeRequest(
run_id="run_4",
node_instance_id="ni_4",
inputs={},
output_dir=str(tmp_path),
)
)
assert response.status == "completed"
assert response.outputs["text"] == "echo"
# ---------------------------------------------------------------------------
# FFmpeg 节点
# ---------------------------------------------------------------------------
def _ffmpeg_request(tmp_path, **overrides) -> InvokeRequest:
"""构造包含默认视频输入与提音参数的调用请求。"""
payload = {
"run_id": "run_1",
"node_instance_id": "ni_1",
"inputs": {"video_uri": str(tmp_path / "input.mp4")},
"params": {"sample_rate": 16000, "channels": 1},
"output_dir": str(tmp_path / "out"),
}
payload.update(overrides)
return InvokeRequest(**payload)
def test_ffmpeg_success(tmp_path, monkeypatch) -> None:
"""验证成功调用会生成 audio.wav 产物。"""
def fake_run(command, **kwargs):
output = Path(command[-1])
output.parent.mkdir(parents=True, exist_ok=True)
output.write_bytes(b"fake wav")
return subprocess.CompletedProcess(command, 0)
monkeypatch.setattr("nodes.ffmpeg.shutil.which", lambda _: "ffmpeg")
monkeypatch.setattr("nodes.ffmpeg.subprocess.run", fake_run)
response = ffmpeg_invoke(_ffmpeg_request(tmp_path))
assert response.status == "completed"
assert Path(response.outputs["audio_uri"]).name == "audio.wav"
def test_ffmpeg_configured_bin(tmp_path, monkeypatch) -> None:
"""验证 FFMPEG_BIN 环境变量优先于 PATH 查找。"""
fake_bin = tmp_path / "ffmpeg.exe"
fake_bin.write_bytes(b"")
monkeypatch.setenv("FFMPEG_BIN", str(fake_bin))
monkeypatch.setattr("nodes.ffmpeg.subprocess.run", lambda *a, **k: subprocess.CompletedProcess([], 0))
response = ffmpeg_invoke(_ffmpeg_request(tmp_path))
assert response.status == "completed"
def test_ffmpeg_bundled_fallback(tmp_path, monkeypatch) -> None:
"""验证无系统 ffmpeg 时回退到 imageio-ffmpeg 内置二进制。"""
monkeypatch.delenv("FFMPEG_BIN", raising=False)
monkeypatch.setattr("nodes.ffmpeg.shutil.which", lambda _: None)
bundled = _ffmpeg_bin()
assert bundled != "ffmpeg"
assert Path(bundled).is_file()
monkeypatch.setattr("nodes.ffmpeg._ffmpeg_bin", lambda: bundled)
def fake_run(command, **kwargs):
Path(command[-1]).parent.mkdir(parents=True, exist_ok=True)
Path(command[-1]).write_bytes(b"wav")
return subprocess.CompletedProcess(command, 0)
monkeypatch.setattr("nodes.ffmpeg.subprocess.run", fake_run)
response = ffmpeg_invoke(_ffmpeg_request(tmp_path))
assert response.status == "completed"
def test_ffmpeg_bundled_import_error(monkeypatch) -> None:
"""验证 imageio-ffmpeg 不可用时最终回退为 "ffmpeg" 字符串。"""
monkeypatch.delenv("FFMPEG_BIN", raising=False)
monkeypatch.setattr("nodes.ffmpeg.shutil.which", lambda _: None)
monkeypatch.setitem(sys.modules, "imageio_ffmpeg", None)
assert _ffmpeg_bin() == "ffmpeg"
def test_ffmpeg_missing_video_uri(tmp_path) -> None:
"""验证缺少 video_uri 时返回失败。"""
response = ffmpeg_invoke(_ffmpeg_request(tmp_path, inputs={}))
assert response.status == "failed"
assert "video_uri" in response.error
def test_ffmpeg_missing_bin(tmp_path, monkeypatch) -> None:
"""验证找不到任何 ffmpeg 时返回明确失败信息。"""
monkeypatch.delenv("FFMPEG_BIN", raising=False)
monkeypatch.setattr("nodes.ffmpeg.shutil.which", lambda _: None)
monkeypatch.setattr("nodes.ffmpeg._bundled_ffmpeg", lambda: None)
response = ffmpeg_invoke(_ffmpeg_request(tmp_path))
assert response.status == "failed"
assert "ffmpeg not found" in response.error
def test_ffmpeg_failure(tmp_path, monkeypatch) -> None:
"""验证 ffmpeg 返回非零退出码时透传 stderr 错误。"""
def fake_run(command, **kwargs):
return subprocess.CompletedProcess(command, 1, stderr="boom")
monkeypatch.setattr("nodes.ffmpeg.shutil.which", lambda _: "ffmpeg")
monkeypatch.setattr("nodes.ffmpeg.subprocess.run", fake_run)
response = ffmpeg_invoke(_ffmpeg_request(tmp_path))
assert response.status == "failed"
assert "boom" in response.error
# ---------------------------------------------------------------------------
# Whisper 节点
# ---------------------------------------------------------------------------
class FakeSegment:
"""模拟 faster-whisper 的分段对象,只提供转写测试需要的字段。"""
def __init__(self, start, end, text):
self.start = start
self.end = end
self.text = text
class FakeWhisperModel:
"""记录构造参数并返回固定分段的假 WhisperModel。"""
instances: list[tuple[tuple, dict]] = []
def __init__(self, *args, **kwargs):
# 记录每次构造参数,测试据此断言 device/compute_type 传递。
FakeWhisperModel.instances.append((args, kwargs))
self.args = args
self.kwargs = kwargs
def transcribe(self, path, **kwargs):
# 返回固定两个分段:一个普通时长,一个跨小时验证时间戳格式。
return (
[
FakeSegment(0, 1, "第一段"),
FakeSegment(3600.5, 3602.25, "第二段"),
],
None,
)
def _install_fake_whisper(monkeypatch, model_class=FakeWhisperModel) -> None:
"""把假 faster_whisper 模块注入 sys.modules,替代真实依赖。"""
fake_module = types.SimpleNamespace(WhisperModel=model_class)
monkeypatch.setitem(sys.modules, "faster_whisper", fake_module)
def _whisper_request(tmp_path, **overrides) -> InvokeRequest:
"""构造默认音频输入与日语参数的调用请求。"""
payload = {
"run_id": "run_1",
"node_instance_id": "ni_1",
"inputs": {"audio_uri": str(tmp_path / "audio.wav")},
"params": {"language": "ja"},
"output_dir": str(tmp_path / "out"),
}
payload.update(overrides)
return InvokeRequest(**payload)
def test_resolve_explicit_param_wins(tmp_path) -> None:
"""验证请求参数中的 model_path 优先级最高,覆盖环境变量与本地候选。"""
local = tmp_path / "model"
local.mkdir()
(local / "model.bin").write_bytes(b"x")
resolved = resolve_model_path(
{"model_path": "/opt/custom-model"},
env={"WHISPER_MODEL_PATH": "/env/model"},
candidates=[local],
)
assert resolved == "/opt/custom-model"
def test_resolve_env_wins_over_local(tmp_path) -> None:
"""验证 WHISPER_MODEL_PATH 环境变量优先于本地候选目录。"""
local = tmp_path / "model"
local.mkdir()
(local / "model.bin").write_bytes(b"x")
resolved = resolve_model_path(
{},
env={"WHISPER_MODEL_PATH": "/env/model"},
candidates=[local],
)
assert resolved == "/env/model"
def test_resolve_local_candidate_used(tmp_path) -> None:
"""验证无参数与环境变量时优先使用含 model.bin 的本地候选目录。"""
local = tmp_path / "model"
local.mkdir()
(local / "model.bin").write_bytes(b"x")
resolved = resolve_model_path({}, env={}, candidates=[local])
assert resolved == str(local)
def test_resolve_incomplete_candidate_skipped(tmp_path) -> None:
"""验证缺少 model.bin 的候选目录被跳过,避免加载残缺模型。"""
empty = tmp_path / "empty"
empty.mkdir()
resolved = resolve_model_path({}, env={}, candidates=[empty])
assert resolved == "large-v2"
def test_resolve_fallback_remote() -> None:
"""验证全部本地候选缺失时回退到远端 large-v2 作为最后兜底。"""
resolved = resolve_model_path({}, env={}, candidates=[])
assert resolved == "large-v2"
def test_format_timestamp() -> None:
"""验证秒数到 SRT 时间戳的格式化结果。"""
assert format_timestamp(0) == "00:00:00,000"
assert format_timestamp(3600.5) == "01:00:00,500"
assert format_timestamp(61.25) == "00:01:01,250"
def test_whisper_success(tmp_path, monkeypatch) -> None:
"""验证成功转写会生成 SRT 并默认使用 auto 设备、float16 计算类型。"""
FakeWhisperModel.instances.clear()
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(_whisper_request(tmp_path))
assert response.status == "completed"
content = Path(response.outputs["srt_uri"]).read_text(encoding="utf-8")
assert "第一段" in content
assert "01:00:00,500 --> 01:00:02,250" in content
_, kwargs = FakeWhisperModel.instances[-1]
assert kwargs["device"] == "auto"
assert kwargs["compute_type"] == "float16"
def test_whisper_pause_flag_stops_between_chunks(tmp_path, monkeypatch) -> None:
"""验证 paused.flag 存在时 whisper 在分块边界中止(批量暂停机制)。
run 根目录(output_dir 的上上级)的 paused.flag 由暂停接口写入,whisper
在每个分块转写前检查;检测到即抛异常,invoke 统一转 failed 响应,调度器
捕获后保持任务 PAUSED。
"""
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
# 暂停信号位于 run 根目录:output_dir 为 runs/run_p/steps/whisper
# 其上上级即 runs/run_p(与调度器/OCR 的目录约定一致)。
(tmp_path / "runs" / "run_p").mkdir(parents=True)
(tmp_path / "runs" / "run_p" / "paused.flag").write_text("", encoding="utf-8")
# 注入两个分块路径,确保进入分块循环并执行至少一次暂停检查。
monkeypatch.setattr(
"nodes.whisper._split_audio",
lambda *args, **kwargs: [tmp_path / "chunk_1.wav", tmp_path / "chunk_2.wav"],
)
response = whisper_invoke(
_whisper_request(
tmp_path,
params={"language": "ja", "chunk_seconds": 60},
output_dir=str(tmp_path / "runs" / "run_p" / "steps" / "whisper"),
)
)
assert response.status == "failed"
assert "被暂停" in response.error
# 暂停时不会写出 SRT 产物。
assert not (tmp_path / "runs" / "run_p" / "steps" / "whisper" / "transcript.srt").exists()
def test_load_cuda_libraries_linux(monkeypatch) -> None:
"""验证 Linux 下进程内预加载 nvidia 动态库,含失败跳过分支。"""
import ctypes
from nodes.whisper import _load_cuda_libraries
# 用真实 nvidia 轮子路径加载,不应抛出异常。
_load_cuda_libraries()
# 模拟加载失败分支:部分库抛 OSError 时应被跳过。
calls: list[str] = []
real_cdll = ctypes.CDLL
def fake_cdll(path):
calls.append(str(path))
if "cudnn" in str(path):
raise OSError("boom")
return real_cdll(path)
monkeypatch.setattr("nodes.whisper.ctypes.CDLL", fake_cdll)
_load_cuda_libraries()
assert calls
def test_load_cuda_libraries_windows(monkeypatch, tmp_path) -> None:
"""验证 Windows 分支通过 add_dll_directory 注册 DLL 搜索目录。"""
import os
import sysconfig
from nodes.whisper import _load_cuda_libraries
site = tmp_path / "site"
(site / "nvidia" / "cublas" / "bin").mkdir(parents=True)
(site / "nvidia" / "cudnn" / "bin").mkdir(parents=True)
monkeypatch.setattr(sysconfig, "get_paths", lambda: {"purelib": str(site)})
added: list[str] = []
# 只注入平台判断与 DLL 目录注册,不改动全局 os.name,避免 pathlib 出错。
monkeypatch.setattr("nodes.whisper._is_windows", lambda: True)
monkeypatch.setattr(os, "add_dll_directory", lambda d: added.append(d), raising=False)
_load_cuda_libraries()
assert any("cublas" in d and d.endswith("bin") for d in added)
def test_whisper_compute_type_override(tmp_path, monkeypatch) -> None:
"""验证请求参数可以覆盖默认计算类型。"""
FakeWhisperModel.instances.clear()
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(
_whisper_request(tmp_path, params={"language": "ja", "compute_type": "int8"})
)
assert response.status == "completed"
_, kwargs = FakeWhisperModel.instances[-1]
assert kwargs["compute_type"] == "int8"
def test_whisper_model_raises(tmp_path, monkeypatch) -> None:
"""验证模型加载失败时返回 failed 与错误信息。"""
class BrokenModel:
def __init__(self, *args, **kwargs):
raise RuntimeError("model load failed")
_install_fake_whisper(monkeypatch, BrokenModel)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(_whisper_request(tmp_path))
assert response.status == "failed"
assert "model load failed" in response.error
def test_whisper_missing_input(tmp_path) -> None:
"""验证缺少 audio_uri 时返回失败。"""
response = whisper_invoke(_whisper_request(tmp_path, inputs={}))
assert response.status == "failed"
def test_whisper_missing_file(tmp_path) -> None:
"""验证音频文件不存在时返回失败。"""
response = whisper_invoke(_whisper_request(tmp_path))
assert response.status == "failed"
assert "audio file not found" in response.error
# ---------------------------------------------------------------------------
# LLM 节点
# ---------------------------------------------------------------------------
class FakeUrlOpenResponse:
"""模拟 urllib 响应对象,提供固定 LLM 译文内容。"""
def __init__(self, content: str) -> None:
# 预编码为 Chat Completions 风格的 JSON 响应体。
self._payload = json.dumps(
{"choices": [{"message": {"content": content}}]}
).encode("utf-8")
def read(self) -> bytes:
return self._payload
def __enter__(self):
return self
def __exit__(self, *args) -> bool:
return False
def _make_srt(tmp_path, count=5) -> Path:
"""生成标准 SRT 测试文件,文本行为"原文字幕N"。"""
lines = []
for index in range(count):
lines.extend(
[
str(index + 1),
f"00:00:{index:02d},000 --> 00:00:{index + 1:02d},000",
f"原文字幕{index + 1}",
"",
]
)
path = tmp_path / "in.srt"
path.write_text("\n".join(lines), encoding="utf-8")
return path
def test_llm_translate_lines_via_fake_api(monkeypatch) -> None:
"""验证通过真实 HTTP 服务器调用 LLM 接口并保持行顺序。"""
class Handler(BaseHTTPRequestHandler):
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", "0"))
self.rfile.read(length)
body = json.dumps(
{
"choices": [
{
"message": {
"content": json.dumps([{"id": i, "text": text} for i, text in enumerate(["译文一", "译文二", "译文三", "译文四", "译文五"], 1)])
}
}
]
}
).encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def log_message(self, format, *args) -> None:
return
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
monkeypatch.setenv(
"LLM_API_BASE",
f"http://127.0.0.1:{server.server_address[1]}/v1/chat/completions",
)
monkeypatch.setenv("LLM_API_KEY", "test-key")
result = translate_lines(
["一", "二", "三", "四", "五"],
{"target_language": "zh-CN"},
)
assert result == ["译文一", "译文二", "译文三", "译文四", "译文五"]
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
def test_llm_translate_lines_empty_input() -> None:
"""验证空行列表直接返回空译文,不发起任何 LLM 请求。
翻译开始前对空输入提前返回:没有字幕行时无需计算批数、也无需调用
接口(避免无谓请求与除零等边界问题),同时保持返回类型为列表。"""
result = translate_lines([], {})
assert result == []
def test_llm_translate_lines_api_error(monkeypatch) -> None:
"""验证 LLM 接口不可用时抛出 URLError。"""
def fail_open(request, timeout):
raise urllib.error.URLError("api down")
monkeypatch.setattr("nodes.llm.urllib.request.urlopen", fail_open)
try:
translate_lines(["一"], {})
raise AssertionError("expected failure")
except urllib.error.URLError:
pass
def test_llm_translate_lines_default_timeout(monkeypatch) -> None:
"""验证未配置超时时使用默认 600 秒。"""
captured = {}
def fake_open(request, timeout):
captured["timeout"] = timeout
return FakeUrlOpenResponse('[{"id": 1, "text": "译文一"}]')
monkeypatch.setattr("nodes.llm.urllib.request.urlopen", fake_open)
monkeypatch.setenv("LLM_API_BASE", "http://fake/v1/chat/completions")
result = translate_lines(["一"], {})
assert result == ["译文一"]
assert captured["timeout"] == 600
def test_llm_translate_lines_env_timeout(monkeypatch) -> None:
"""验证 LLM_TIMEOUT_SECONDS 环境变量可覆盖超时。"""
captured = {}
def fake_open(request, timeout):
captured["timeout"] = timeout
return FakeUrlOpenResponse('[{"id": 1, "text": "译文一"}]')
monkeypatch.setattr("nodes.llm.urllib.request.urlopen", fake_open)
monkeypatch.setenv("LLM_API_BASE", "http://fake/v1/chat/completions")
monkeypatch.setenv("LLM_TIMEOUT_SECONDS", "45")
translate_lines(["一"], {})
assert captured["timeout"] == 45
def test_llm_invoke_success(tmp_path, monkeypatch) -> None:
"""验证成功调用会把译文回填到 SRT 并输出 cn.srt。"""
source = _make_srt(tmp_path)
def fake_translate(lines, params):
return [f"译文{i + 1}" for i in range(len(lines))]
monkeypatch.setattr("nodes.llm.translate_lines", fake_translate)
response = llm_invoke(
InvokeRequest(
run_id="run_1",
node_instance_id="ni_1",
inputs={"srt_uri": str(source)},
params={"target_language": "zh-CN"},
output_dir=str(tmp_path / "out"),
)
)
assert response.status == "completed"
content = Path(response.outputs["cn_srt_uri"]).read_text(encoding="utf-8")
assert "译文1" in content
def test_llm_invoke_rejects_short_translation(tmp_path, monkeypatch) -> None:
"""防御性校验:译文条数不足时失败,不用空行掩盖不完整结果。"""
source = _make_srt(tmp_path, count=3)
monkeypatch.setattr(
"nodes.llm.translate_lines",
lambda lines, params: ["only one"],
)
response = llm_invoke(
InvokeRequest(
run_id="run_2",
node_instance_id="ni_2",
inputs={"srt_uri": str(source)},
output_dir=str(tmp_path / "out2"),
)
)
assert response.status == "failed"
def test_llm_invoke_missing_input(tmp_path) -> None:
"""验证缺少 srt_uri 时返回失败。"""
response = llm_invoke(
InvokeRequest(
run_id="run_3",
node_instance_id="ni_3",
inputs={},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
def test_llm_invoke_missing_file(tmp_path) -> None:
"""验证 SRT 文件不存在时返回失败。"""
response = llm_invoke(
InvokeRequest(
run_id="run_4",
node_instance_id="ni_4",
inputs={"srt_uri": str(tmp_path / "missing.srt")},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
# ---------------------------------------------------------------------------
# ASS 节点
# ---------------------------------------------------------------------------
SAMPLE_SRT = """
1
00:00:01,000 --> 00:00:03,000
第一行
第二行
2
00:00:04,000 --> 00:00:06,000
第三行
"""
def test_ass_parse_and_write(tmp_path) -> None:
"""验证多行字幕会被解析并通过左右眼样式写出。"""
entries = parse_srt(SAMPLE_SRT)
assert len(entries) == 2
assert entries[0][2] == r"第一行\N第二行"
output = tmp_path / "out.ass"
write_ass(entries, output, "3840x1920")
content = output.read_text(encoding="utf-8")
assert "PlayResX: 3840" in content
assert "PlayResY: 1920" in content
assert "LeftEye" in content
assert "RightEye" in content
assert r"第一行\N第二行" in content
def test_ass_top_aligned_and_translucent(tmp_path) -> None:
"""验证字幕默认渲染到顶部安全区且文字/描边带透明度。
B-1:对齐 an8(顶部居中),MarginV 取顶部安全边距默认 7002026-09 起,
字幕置于视线自然可读位置;旧默认 120 落在画面最顶需抬头观看);
透明度:文字填充 &HB3FFFFFF(约 70% 透明),描边 &H80000000(半透明黑),
而左右眼水平相对位置一致保持零视差(A-1,字幕在屏幕平面)。"""
entries = parse_srt(SAMPLE_SRT)
output = tmp_path / "out.ass"
write_ass(entries, output, "3840x1920")
content = output.read_text(encoding="utf-8")
# 顶部对齐 an8、顶部安全边距默认 700。
assert r"{\an8}" in content
assert r"{\an2}" not in content
# 样式行:MarginV=700(顶部安全区),填充 &HB3FFFFFF,描边 &H80000000。
assert "&HB3FFFFFF" in content
assert "&H80000000" in content
assert ",0,0,0,0,50,100,0,0,1,4,0,8,50,1920,700,1" in content
# 左右眼两行文本一致(水平相对位置相同 → 零视差/A-1)。
assert content.count(r"第一行\N第二行") == 2
def test_ass_write_custom_margin_top(tmp_path) -> None:
"""验证可通过 margin_top 覆盖顶部安全边距(不同分辨率下微调)。"""
entries = parse_srt(SAMPLE_SRT)
output = tmp_path / "out.ass"
write_ass(entries, output, "1920x1080", margin_top=200)
content = output.read_text(encoding="utf-8")
# 1920 宽:每半幅 960MarginV 用自定义 200。
assert ",0,0,0,0,50,100,0,0,1,4,0,8,50,960,200,1" in content
def test_ass_invoke_reads_margin_top(tmp_path) -> None:
"""验证 invoke 从请求参数读取 margin_top 并反映到产物样式。"""
source = tmp_path / "in.srt"
source.write_text(SAMPLE_SRT, encoding="utf-8")
response = ass_invoke(
InvokeRequest(
run_id="run_mt",
node_instance_id="ni_mt",
inputs={"cn_srt_uri": str(source)},
params={"resolution": "1920x1080", "margin_top": 90},
output_dir=str(tmp_path / "out"),
)
)
assert response.status == "completed"
content = Path(response.outputs["ass_uri"]).read_text(encoding="utf-8")
assert "an8" in content
assert ",0,0,0,0,50,100,0,0,1,4,0,8,50,960,90,1" in content
def test_ass_invoke_default_margin_top_is_700(tmp_path) -> None:
"""验证 invoke 未显式传 margin_top 时默认生成 MarginV=700 的样式。
保证之后生成的字幕默认都落在顶部安全区下方(2026-09 起),
不再回退到旧的 120(需抬头观看)。"""
source = tmp_path / "in.srt"
source.write_text(SAMPLE_SRT, encoding="utf-8")
response = ass_invoke(
InvokeRequest(
run_id="run_def700",
node_instance_id="ni_def700",
inputs={"cn_srt_uri": str(source)},
params={"resolution": "1920x1080"},
output_dir=str(tmp_path / "out"),
)
)
assert response.status == "completed"
content = Path(response.outputs["ass_uri"]).read_text(encoding="utf-8")
assert "an8" in content
assert ",0,0,0,0,50,100,0,0,1,4,0,8,50,960,700,1" in content
def test_ass_parse_malformed(tmp_path) -> None:
"""验证畸形 SRT 不会抛出异常且返回空条目或忽略坏行。"""
source = tmp_path / "bad.srt"
source.write_text("1\nnot a time line\n", encoding="utf-8")
assert parse_srt(source.read_text(encoding="utf-8")) == []
source.write_text("1", encoding="utf-8")
assert parse_srt(source.read_text(encoding="utf-8")) == []
def test_ass_invoke_success(tmp_path) -> None:
"""验证成功调用会按指定分辨率输出 ASS 产物。"""
source = tmp_path / "in.srt"
source.write_text(SAMPLE_SRT, encoding="utf-8")
response = ass_invoke(
InvokeRequest(
run_id="run_1",
node_instance_id="ni_1",
inputs={"cn_srt_uri": str(source)},
params={"resolution": "1920x1080"},
output_dir=str(tmp_path / "out"),
)
)
assert response.status == "completed"
assert "PlayResX: 1920" in Path(response.outputs["ass_uri"]).read_text(encoding="utf-8")
def test_ass_invoke_missing_input(tmp_path) -> None:
"""验证缺少 cn_srt_uri 时返回失败。"""
response = ass_invoke(
InvokeRequest(
run_id="run_2",
node_instance_id="ni_2",
inputs={},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
def test_ass_invoke_missing_file(tmp_path) -> None:
"""验证 SRT 文件不存在时返回失败。"""
response = ass_invoke(
InvokeRequest(
run_id="run_3",
node_instance_id="ni_3",
inputs={"cn_srt_uri": str(tmp_path / "missing.srt")},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
def test_resolve_bare_name_found(tmp_path) -> None:
"""验证裸模型名会在候选目录父级下按名解析到本地模型。"""
named = tmp_path / "zh-ct2"
named.mkdir()
(named / "model.bin").write_bytes(b"x")
# 候选目录取父目录的兄弟布局:candidates 首项父级即模型根目录。
candidates = [tmp_path / "models"]
resolved = resolve_model_path({"model_path": "zh-ct2"}, env={}, candidates=candidates)
assert resolved == str(named)
def test_resolve_bare_name_missing(tmp_path) -> None:
"""验证裸模型名在本地不存在时原样返回,交由 faster-whisper 处理。"""
resolved = resolve_model_path(
{"model_path": "no-such-model"}, env={}, candidates=[tmp_path / "models"]
)
assert resolved == "no-such-model"
def test_whisper_task_translate(tmp_path, monkeypatch) -> None:
"""验证 task=translate 参数会传递给 faster-whisper 的 transcribe。"""
class TaskRecorderModel:
def __init__(self, *args, **kwargs):
self.transcribe_kwargs = None
def transcribe(self, path, **kwargs):
self.transcribe_kwargs = kwargs
return (
[FakeSegment(0, 1, "中文直出")],
None,
)
recorder = TaskRecorderModel()
monkeypatch.setitem(sys.modules, "faster_whisper", types.SimpleNamespace(WhisperModel=lambda *a, **k: recorder))
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(
_whisper_request(tmp_path, params={"language": "ja", "task": "translate"})
)
assert response.status == "completed"
assert recorder.transcribe_kwargs["task"] == "translate"
assert "中文直出" in Path(response.outputs["srt_uri"]).read_text(encoding="utf-8")
def test_whisper_condition_on_previous_text(tmp_path, monkeypatch) -> None:
"""验证长音频参数 condition_on_previous_text 可配置并默认关闭。"""
captured = {}
class CondRecorderModel:
def __init__(self, *args, **kwargs):
pass
def transcribe(self, path, **kwargs):
captured["condition_on_previous_text"] = kwargs.get("condition_on_previous_text")
return ([FakeSegment(0, 1, "ok")], None)
monkeypatch.setitem(
sys.modules,
"faster_whisper",
types.SimpleNamespace(WhisperModel=lambda *a, **k: CondRecorderModel()),
)
_make_wav(tmp_path / "audio.wav", 5)
# 默认 False(长音频稳定);显式传 True 可开启。
whisper_invoke(_whisper_request(tmp_path))
assert captured["condition_on_previous_text"] is False
whisper_invoke(_whisper_request(tmp_path, params={"condition_on_previous_text": True}))
assert captured["condition_on_previous_text"] is True
def test_split_audio_disabled_or_no_ffmpeg(tmp_path, monkeypatch) -> None:
"""验证 chunk_seconds<=0 或缺少 ffmpeg 时回退整段,不调用切块。"""
from nodes.whisper import _split_audio
audio = _make_wav(tmp_path / "in.wav", 5)
called = []
def fake_run(*args, **kwargs):
called.append(args)
return subprocess.CompletedProcess([], 0)
monkeypatch.setattr("nodes.whisper.subprocess.run", fake_run)
# chunk_seconds<=0:直接返回整段,不执行 ffmpeg。
assert _split_audio(audio, tmp_path, 0, "ffmpeg") == [audio]
# 缺少 ffmpeg:直接返回整段。
assert _split_audio(audio, tmp_path, 600, None) == [audio]
assert called == []
def test_split_audio_failure_fallback(tmp_path, monkeypatch) -> None:
"""验证 ffmpeg 切块失败时回退整段单次转写。"""
from nodes.whisper import _split_audio
audio = _make_wav(tmp_path / "in.wav", 5)
monkeypatch.setattr(
"nodes.whisper.subprocess.run",
lambda *a, **k: subprocess.CompletedProcess([], 1, stderr="boom"),
)
assert _split_audio(audio, tmp_path, 600, "ffmpeg") == [audio]
def test_split_audio_success_and_empty(tmp_path, monkeypatch) -> None:
"""验证切块成功返回块列表;产出为空时回退整段。"""
from nodes.whisper import _split_audio
audio = _make_wav(tmp_path / "in.wav", 5)
def fake_run_success(command, **kwargs):
# 模拟 ffmpeg 产出两个真实的块 WAV。
pattern = command[-1]
for name in ("chunk_000.wav", "chunk_001.wav"):
_make_wav(tmp_path / "chunks" / name, 2)
return subprocess.CompletedProcess(command, 0)
monkeypatch.setattr("nodes.whisper.subprocess.run", fake_run_success)
chunks = _split_audio(audio, tmp_path, 600, "ffmpeg")
assert len(chunks) == 2
assert chunks[0].name == "chunk_000.wav"
# 切块成功但没有产出文件时回退整段。
monkeypatch.setattr(
"nodes.whisper.subprocess.run",
lambda *a, **k: subprocess.CompletedProcess([], 0),
)
assert _split_audio(audio, tmp_path / "other", 600, "ffmpeg") == [audio]
def test_whisper_chunked_transcription_merges_offsets(tmp_path, monkeypatch) -> None:
"""验证分块转写用真实 WAV 块,第二块时间轴按实际时长偏移合并到同一 SRT。"""
from nodes.whisper import _split_audio
chunk_dir = tmp_path / "out" / "chunks"
chunk_dir.mkdir(parents=True)
# 真实 WAV 块(各 60s),偏移按实际时长累积为 60s。
chunk1 = _make_wav(chunk_dir / "chunk_000.wav", 60)
chunk2 = _make_wav(chunk_dir / "chunk_001.wav", 60)
monkeypatch.setattr("nodes.whisper._split_audio", lambda a, o, c, f: [chunk1, chunk2])
FakeWhisperModel.instances.clear()
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(_whisper_request(tmp_path, params={"chunk_seconds": 60}))
assert response.status == "completed"
content = Path(response.outputs["srt_uri"]).read_text(encoding="utf-8")
# 序号从 1 开始连续递增;第一块无偏移,第二块偏移 60 秒。
assert content.startswith("1\n")
assert "3\n00:01:00,000" in content
assert "00:00:00,000 --> 00:00:01,000" in content
assert "01:00:00,500 --> 01:00:02,250" in content
assert "00:01:00,000 --> 00:01:01,000" in content
assert "01:01:00,500 --> 01:01:02,250" in content
def test_whisper_chunk_disabled_single_call(tmp_path, monkeypatch) -> None:
"""验证 chunk_seconds=0 时单次调用、不切块。"""
FakeWhisperModel.instances.clear()
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(_whisper_request(tmp_path, params={"chunk_seconds": 0}))
assert response.status == "completed"
# 单次调用:init 只记录一次实例。
assert len(FakeWhisperModel.instances) == 1
def test_whisper_vad_filter_default_off(tmp_path, monkeypatch) -> None:
"""验证 vad_filter 默认开启,可显式关闭。"""
captured = {}
class VadRecorderModel:
def __init__(self, *args, **kwargs):
pass
def transcribe(self, path, **kwargs):
captured["vad_filter"] = kwargs.get("vad_filter")
return ([FakeSegment(0, 1, "ok")], None)
monkeypatch.setitem(
sys.modules,
"faster_whisper",
types.SimpleNamespace(WhisperModel=lambda *a, **k: VadRecorderModel()),
)
_make_wav(tmp_path / "audio.wav", 5)
whisper_invoke(_whisper_request(tmp_path))
assert captured["vad_filter"] is True
whisper_invoke(_whisper_request(tmp_path, params={"vad_filter": False}))
assert captured["vad_filter"] is False
def _make_wav(path, seconds, rate=16000) -> Path:
"""生成指定时长的 16kHz 单声道 16bit 静音 WAV。"""
import wave
with wave.open(str(path), "wb") as wav:
wav.setnchannels(1)
wav.setsampwidth(2)
wav.setframerate(rate)
wav.writeframes(b"\x00\x00" * int(rate * seconds))
return path
def test_whisper_chunk_offset_uses_actual_duration(tmp_path, monkeypatch) -> None:
"""验证分块偏移按 WAV 实际时长累积(1s+2s 块 → 第二块偏移 1s 而非块长 60s)。"""
chunk_dir = tmp_path / "out" / "chunks"
chunk_dir.mkdir(parents=True)
chunk1 = _make_wav(chunk_dir / "chunk_000.wav", 1)
chunk2 = _make_wav(chunk_dir / "chunk_001.wav", 2)
monkeypatch.setattr("nodes.whisper._split_audio", lambda a, o, c, f: [chunk1, chunk2])
FakeWhisperModel.instances.clear()
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(_whisper_request(tmp_path, params={"chunk_seconds": 60}))
assert response.status == "completed"
content = Path(response.outputs["srt_uri"]).read_text(encoding="utf-8")
# 第二块偏移 = 第一块实际时长 1s(若用块长假设则会是 60s → 00:01:00)。
assert "00:00:01,000 --> 00:00:02,000" in content
assert "00:01:00,000 --> 00:01:01,000" not in content
def test_whisper_segment_logs_full_video_time(caplog, tmp_path, monkeypatch) -> None:
"""验证每条分段日志包含编号与完整视频角度(含分块偏移)的时间范围。"""
chunk_dir = tmp_path / "out" / "chunks"
chunk_dir.mkdir(parents=True)
chunk1 = _make_wav(chunk_dir / "chunk_000.wav", 60)
chunk2 = _make_wav(chunk_dir / "chunk_001.wav", 60)
monkeypatch.setattr("nodes.whisper._split_audio", lambda a, o, c, f: [chunk1, chunk2])
FakeWhisperModel.instances.clear()
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
with caplog.at_level("INFO", logger="vrsub.whisper"):
response = whisper_invoke(_whisper_request(tmp_path, params={"chunk_seconds": 60}))
assert response.status == "completed"
seg_logs = [r.message for r in caplog.records if r.message.startswith("分段 #")]
# 第一块:0s 起;第二块:偏移 60s(第一块实际时长)。
assert any("分段 #1: 00:00:00,000 --> 00:00:01,000" in m for m in seg_logs)
assert any(m.startswith("分段 #3: 00:01:00,000") for m in seg_logs)
def test_whisper_auto_vad_fallback_on_error(tmp_path, monkeypatch) -> None:
"""自动 VAD 分析抛异常时,whisper 回退默认参数正常转写(防御性)。"""
import nodes.whisper as whisper_mod
def _boom(*args, **kwargs):
raise RuntimeError("auto vad analysis failed")
# 使 nodes.vad_profiler.vad_parameters_for_audio 抛异常(whisper.py 函数内 import)。
monkeypatch.setattr(
"nodes.vad_profiler.vad_parameters_for_audio", _boom, raising=False,
)
_install_fake_whisper(monkeypatch)
_make_wav(tmp_path / "audio.wav", 5)
response = whisper_invoke(
_whisper_request(tmp_path, params={"language": "ja", "vad_filter": True, "sample_rate": 16000})
)
assert response.status == "completed"
# 回退默认:transcribe 收到 vad_parameters=None。
_, kwargs = FakeWhisperModel.instances[-1]
assert kwargs.get("vad_parameters") is None
def test_wav_duration_fallback_on_invalid_file(tmp_path) -> None:
"""验证读取时长时,损坏/缺失文件回退 fallback 值(真实非法文件,非占位字节)。"""
from nodes.whisper import _wav_duration_seconds
bad = tmp_path / "bad.wav"
bad.write_text("this is not a wav file", encoding="utf-8")
assert _wav_duration_seconds(bad, 60.0) == 60.0
missing = tmp_path / "missing.wav"
assert _wav_duration_seconds(missing, 60.0) == 60.0
# ---------------------------------------------------------------------------
# VLM OCR 节点(直接请求 Ollama /api/generate,流式)
# ---------------------------------------------------------------------------
# 真实的最小 PNG(1x1 像素,合法文件),用于构造真实图片输入。
_MIN_PNG_BASE64 = (
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8"
"z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
)
class FakeChatResponse:
"""模拟 Ollama /api/generate 流式响应:readline() 逐行返回 JSON 块。
done=False(默认)时仅一行内容,随后 readline 返回空表示流结束;
done=True 时末尾追加一行带 done 标记的结束块。
"""
def __init__(self, content: str, done: bool = False) -> None:
self._lines = [
json.dumps({"message": {"role": "assistant", "content": content}}).encode()
]
if done:
self._lines.append(
json.dumps({"message": {"role": "assistant", "content": ""}, "done": True}).encode()
)
def readline(self) -> bytes:
return self._lines.pop(0) if self._lines else b""
def __enter__(self):
return self
def __exit__(self, *args) -> bool:
return False
def _real_png(tmp_path) -> Path:
"""生成真实 PNG 图片文件(解码自合法 base64)。"""
image = tmp_path / "frame.png"
image.write_bytes(base64.b64decode(_MIN_PNG_BASE64))
return image
def test_vlm_success(tmp_path, monkeypatch) -> None:
"""验证真实图片经流式读取后输出清洗过的文字与 ocr.txt 产物。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
captured = {}
def fake_urlopen(request, timeout):
captured["timeout"] = timeout
captured["url"] = request.full_url
return FakeChatResponse("HELLO WORLD 123\n```markdown\n```\n```\n")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
response = vlm_invoke(
InvokeRequest(
run_id="run_1",
node_instance_id="",
inputs={"image_uri": str(image)},
params={"model": "glm-ocr:latest"},
output_dir=str(tmp_path / "out"),
)
)
assert response.status == "completed", response.error
assert response.outputs["text"] == "HELLO WORLD 123"
assert captured["url"].endswith("/api/chat")
assert captured["timeout"] == 5
content = Path(response.outputs["text_uri"]).read_text(encoding="utf-8")
assert "HELLO WORLD 123" in content
def test_vlm_missing_input_and_file(tmp_path) -> None:
"""验证缺少 image_uri 或图片不存在时返回失败。"""
from nodes.vlm import invoke as vlm_invoke
assert vlm_invoke(
InvokeRequest(run_id="r", node_instance_id="", inputs={}, output_dir=str(tmp_path))
).status == "failed"
assert vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(tmp_path / "missing.png")},
output_dir=str(tmp_path),
)
).status == "failed"
def test_vlm_network_error(tmp_path, monkeypatch) -> None:
"""验证 Ollama 服务不可达时返回 failed。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
def fail_open(request, timeout):
raise urllib.error.URLError("ollama down")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fail_open)
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
assert "ollama down" in response.error
def test_vlm_response_format_error(tmp_path, monkeypatch) -> None:
"""验证流式块既无 response 也无 done 标记时(格式错误)返回 failed。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
class BadResponse:
def readline(self) -> bytes:
return b'{"foo": 1}'
def __enter__(self):
return self
def __exit__(self, *args) -> bool:
return False
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", lambda *a, **k: BadResponse())
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
def test_vlm_clean_ocr_text() -> None:
"""验证清洗逻辑会剔除 markdown 围栏与空行。"""
from nodes.vlm import _clean_ocr_text
cleaned = _clean_ocr_text("第一行\n```markdown\n```\n\n第二行\n```")
assert cleaned == "第一行\n第二行"
def test_vlm_options_in_body(tmp_path, monkeypatch) -> None:
"""验证请求体携带采样选项(temperature=0/repeat_penalty/num_predict),可参数覆盖。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
captured = {}
def fake_urlopen(request, timeout):
import json as _json
captured["body"] = _json.loads(request.data.decode("utf-8"))
return FakeChatResponse("SUB 001")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
params={"repeat_penalty": 1.3, "num_predict": 128},
output_dir=str(tmp_path),
)
)
assert captured["body"]["options"] == {
"temperature": 0.3,
"repeat_penalty": 1.3,
"num_predict": 128,
}
def test_vlm_call_structure_system_prompt_and_stop(tmp_path, monkeypatch) -> None:
"""验证 /api/chat 调用结构:system 承载指令、user 只携带图片、stop=</attached_files>。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
captured = {}
def fake_urlopen(request, timeout):
import json as _json
captured["body"] = _json.loads(request.data.decode("utf-8"))
return FakeChatResponse("SUB 001")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
params={"prompt": "识别图片中的所有文字,原样输出。"},
output_dir=str(tmp_path),
)
)
body = captured["body"]
assert body["stop"] == ["\n", "\n答", "答"]
assert body["stream"] is True
assert body["messages"][0]["role"] == "system"
assert "识别图片" in body["messages"][0]["content"]
assert body["messages"][1]["role"] == "user"
assert body["messages"][1]["content"] == ""
assert len(body["messages"][1]["images"]) == 1
def test_vlm_stream_stop_sequence_truncation(tmp_path, monkeypatch) -> None:
"""验证流式读取命中终止序列(换行+答,\n答)即停止并截断该标记。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
def fake_urlopen(request, timeout):
# 模型输出先给出真实文本,随后进入“答:”式重复循环并输出终止序列。
return FakeChatResponse("SUB 001\n答:重复循环垃圾")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
output_dir=str(tmp_path),
)
)
assert response.status == "completed", response.error
# 终止序列及其后内容必须被截掉,只保留终止前的真实识别文本。
assert response.outputs["text"] == "SUB 001"
def test_vlm_stream_done_with_message(tmp_path, monkeypatch) -> None:
"""验证带 done 标记的流式结束块触发停止,结果正常返回。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
def fake_urlopen(request, timeout):
return FakeChatResponse("SUB 001", done=True)
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
output_dir=str(tmp_path),
)
)
assert response.status == "completed", response.error
assert response.outputs["text"] == "SUB 001"
def test_vlm_stream_done_only_end(tmp_path, monkeypatch) -> None:
"""验证仅含 done 标记(无 response)的行视为流结束,空结果正常完成。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
class DoneOnlyResponse:
def readline(self) -> bytes:
if not self._consumed:
self._consumed = True
return b'{"done": true}'
return b""
def __enter__(self):
self._consumed = False
return self
def __exit__(self, *args) -> bool:
return False
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", lambda *a, **k: DoneOnlyResponse())
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
output_dir=str(tmp_path),
)
)
assert response.status == "completed", response.error
assert response.outputs["text"] == ""
def test_vlm_stream_deadline_timeout(tmp_path, monkeypatch) -> None:
"""验证整体 5 秒截止:流式读取超过 deadline 立即终止并返回 failed。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
# 第一次调用计算 deadline(100+5=105),第二次调用已越过截止(200>=105)。
monotonic_values = iter([100.0, 200.0])
monkeypatch.setattr("nodes.vlm.time.monotonic", lambda: next(monotonic_values))
def fake_urlopen(request, timeout):
return FakeChatResponse("SUB 001")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
params={"timeout_seconds": 5},
output_dir=str(tmp_path),
)
)
assert response.status == "failed"
assert "timed out" in response.error
def test_vlm_extract_gettext() -> None:
"""验证从模型输出中提取 <gettext></gettext> 标签内容的各种形态。"""
from nodes.vlm import _extract_gettext
# 正常:标签包裹的内容被提取。
assert _extract_gettext("<gettext>还有没有什么困扰 或者奇怪的地方吗</gettext>") == (
"还有没有什么困扰 或者奇怪的地方吗"
)
# 空标签:返回空字符串。
assert _extract_gettext("前缀<gettext></gettext>后缀") == ""
# 多个标签(重复循环):只取第一个。
assert _extract_gettext("<gettext>SUB 001</gettext><gettext>SUB 001</gettext>") == "SUB 001"
# 跨行内容:DOTALL 让 . 匹配换行。
assert _extract_gettext("<gettext>第一行\n第二行</gettext>") == "第一行\n第二行"
# 未按格式输出(无标签):回退原始文本,保持旧行为。
assert _extract_gettext("没有标签的裸文本") == "没有标签的裸文本"
def test_vlm_gettext_in_invoke(tmp_path, monkeypatch) -> None:
"""验证整条调用链:模型返回 <gettext> 包裹内容时,产物只含标签内文本。"""
from nodes.vlm import invoke as vlm_invoke
image = _real_png(tmp_path)
def fake_urlopen(request, timeout):
return FakeChatResponse("<gettext>SUB 001</gettext> 围栏垃圾```")
monkeypatch.setattr("nodes.vlm.urllib.request.urlopen", fake_urlopen)
response = vlm_invoke(
InvokeRequest(
run_id="r", node_instance_id="",
inputs={"image_uri": str(image)},
output_dir=str(tmp_path),
)
)
assert response.status == "completed", response.error
assert response.outputs["text"] == "SUB 001"
def test_vlm_truncate_at_stop() -> None:
"""验证多个终止序列的截断:取最先命中位置,未命中原样返回。"""
from nodes.vlm import _truncate_at_stop
# 命中 "\n答"(位置更靠前)。
assert _truncate_at_stop("SUB 001\n答:重复") == "SUB 001"
# 未命中 "\n答" 但命中单个 "答"。
assert _truncate_at_stop("SUB 001 答") == "SUB 001 "
# 多个序列均命中:取最早出现的位置("答" 在 "答:" 之前)。
assert _truncate_at_stop("SUB 001答\n答:循环") == "SUB 001"
# 未命中任何序列:原样返回。
assert _truncate_at_stop("还有没有什么困扰 或者奇怪的地方吗") == (
"还有没有什么困扰 或者奇怪的地方吗"
)
# ---------------------------------------------------------------------------
# decode_full:无 VAD 整段解码 + 日语幻觉清洗(TDD)
# ---------------------------------------------------------------------------
def test_whisper_decode_full_disables_vad(tmp_path, monkeypatch) -> None:
"""decode_full=true 时 transcribe 收到 vad_filter=False 且不触发自动 VAD。"""
captured = {}
class FullModel:
def __init__(self, *args, **kwargs):
pass
def transcribe(self, path, **kwargs):
captured["vad_filter"] = kwargs.get("vad_filter")
captured["vad_parameters"] = kwargs.get("vad_parameters")
return ([FakeSegment(0, 1, "ok")], None)
monkeypatch.setitem(
sys.modules, "faster_whisper", types.SimpleNamespace(WhisperModel=lambda *a, **k: FullModel())
)
_make_wav(tmp_path / "audio.wav", 5)
# 默认 vad_filter=true(保持现状),显式 decode_full=true 强制无 VAD 整段解码。
resp = whisper_invoke(
_whisper_request(tmp_path, params={"language": "ja", "decode_full": True})
)
assert resp.status == "completed"
assert captured["vad_filter"] is False
assert captured["vad_parameters"] is None
def test_whisper_decode_full_cleanup_jp_hallucination(tmp_path, monkeypatch) -> None:
"""decode_full=true 时产物 SRT 中日文长时寒暄幻觉被**整条删除**(连带时间戳)。"""
class HallucModel:
def __init__(self, *args, **kwargs):
pass
def transcribe(self, path, **kwargs):
# 模拟无 VAD 解码:正常句 + 一条 30s '晚安'幻觉占位 + 一条 2s 真实晚安。
return (
[
FakeSegment(0, 2, "気持ちいい"),
FakeSegment(10, 40, "おやすみなさい"), # 30s 幻觉
FakeSegment(45, 47, "おやすみなさい"), # 2s 真实
],
None,
)
monkeypatch.setitem(
sys.modules, "faster_whisper", types.SimpleNamespace(WhisperModel=lambda *a, **k: HallucModel())
)
_make_wav(tmp_path / "audio.wav", 5)
resp = whisper_invoke(
_whisper_request(tmp_path, params={"language": "ja", "decode_full": True})
)
assert resp.status == "completed"
content = Path(resp.outputs["srt_uri"]).read_text(encoding="utf-8")
# 30s 幻觉整条删除(时间轴 10-40s 不出现);2s 真实晚安与正常句保留。
assert "気持ちいい" in content
assert content.count("おやすみなさい") == 1
assert "00:00:10,000 --> 00:00:40,000" not in content
# 不残留 '-' 占位(时间轴里的 '-' 是 SRT 合法分隔符,只检查文本行)。
text_lines = [
l for l in content.splitlines()
if l.strip() and not l.strip().isdigit() and "-->" not in l
]
assert all(t != "-" for t in text_lines)
def test_whisper_decode_full_filters_short_moan(tmp_path, monkeypatch) -> None:
"""decode_full=true 时纯呻吟碎片整条删除,真实短词保留(2026-09 决策)。"""
class MoanModel:
def __init__(self, *args, **kwargs):
pass
def transcribe(self, path, **kwargs):
# 混合:正常句 + 纯呻吟碎片 + 真实短词(都应保留)。
return (
[
FakeSegment(0, 2, "気持ちいい"), # 正常
FakeSegment(2, 3, "あ〜"), # 纯呻吟 → 删
FakeSegment(3, 4, "ん?"), # 纯呻吟 → 删
FakeSegment(4, 5, "そこ"), # 真实短词 → 保留
FakeSegment(5, 6, "やばい"), # 真实短词 → 保留
],
None,
)
monkeypatch.setitem(
sys.modules, "faster_whisper", types.SimpleNamespace(WhisperModel=lambda *a, **k: MoanModel())
)
_make_wav(tmp_path / "audio.wav", 5)
resp = whisper_invoke(
_whisper_request(tmp_path, params={"language": "ja", "decode_full": True})
)
assert resp.status == "completed"
content = Path(resp.outputs["srt_uri"]).read_text(encoding="utf-8")
# 纯呻吟碎片被整条删除;真实短词与正常句保留。
assert "あ〜" not in content
assert "ん?" not in content
assert "気持ちいい" in content
assert "そこ" in content
assert "やばい" in content
assert content.count("-->") == 3 # 5 段 - 2 段呻吟 = 3 条
def test_whisper_decode_full_short_moan_can_disable(tmp_path, monkeypatch) -> None:
"""short_moan_max_chars=0 时关闭短呻吟过滤,所有内容原样保留。"""
class MoanModel2:
def __init__(self, *args, **kwargs):
pass
def transcribe(self, path, **kwargs):
return ([FakeSegment(0, 1, "あ〜"), FakeSegment(1, 2, "そこ")], None)
monkeypatch.setitem(
sys.modules, "faster_whisper", types.SimpleNamespace(WhisperModel=lambda *a, **k: MoanModel2())
)
_make_wav(tmp_path / "audio.wav", 5)
resp = whisper_invoke(
_whisper_request(
tmp_path,
params={"language": "ja", "decode_full": True, "short_moan_max_chars": 0},
)
)
assert resp.status == "completed"
content = Path(resp.outputs["srt_uri"]).read_text(encoding="utf-8")
assert "あ〜" in content
assert "そこ" in content
assert content.count("-->") == 2