from __future__ import annotations
|
|
import json
|
import os
|
import subprocess
|
import sys
|
import tempfile
|
import types
|
import unittest
|
from pathlib import Path
|
from unittest.mock import patch
|
|
|
PROJECT_DEV = Path(__file__).resolve().parents[1]
|
if str(PROJECT_DEV) not in sys.path:
|
sys.path.insert(0, str(PROJECT_DEV))
|
|
import transcribe_media
|
|
|
class FakeSegment:
|
def __init__(self, start: float, end: float, text: str) -> None:
|
self.start = start
|
self.end = end
|
self.text = text
|
|
|
class FakeModel:
|
def __init__(self) -> None:
|
self.calls: list[tuple[str, dict[str, object]]] = []
|
|
def transcribe(self, audio: str, **options: object):
|
self.calls.append((audio, options))
|
info = types.SimpleNamespace(language="zh", language_probability=0.99)
|
segments = iter(
|
[
|
FakeSegment(3.12, 8.64, " 大家好。 "),
|
FakeSegment(9.0, 12.5, "开始测试。"),
|
]
|
)
|
return segments, info
|
|
|
class ScriptedModel:
|
def __init__(self, responses, *, fail_on_call: int | None = None) -> None:
|
self.responses = responses
|
self.fail_on_call = fail_on_call
|
self.calls: list[tuple[str, dict[str, object]]] = []
|
|
def transcribe(self, audio: str, **options: object):
|
self.calls.append((audio, options))
|
call_number = len(self.calls)
|
if self.fail_on_call == call_number:
|
raise RuntimeError("synthetic GPU failure")
|
info = types.SimpleNamespace(language="zh", language_probability=0.98)
|
return iter(self.responses[call_number - 1]), info
|
|
|
class RecordingRunner:
|
def __init__(
|
self,
|
*,
|
has_audio: bool = True,
|
ffmpeg_ok: bool = True,
|
audio_duration: float = 60.0,
|
) -> None:
|
self.has_audio = has_audio
|
self.ffmpeg_ok = ffmpeg_ok
|
self.audio_duration = audio_duration
|
self.commands: list[list[str]] = []
|
|
def __call__(self, command):
|
command = list(command)
|
self.commands.append(command)
|
executable = Path(command[0]).name.lower()
|
if executable.startswith("ffprobe"):
|
if "format=duration" in command:
|
stdout = f"{self.audio_duration}\n"
|
else:
|
stdout = "0\n" if self.has_audio else ""
|
return subprocess.CompletedProcess(command, 0, stdout, "")
|
if executable.startswith("ffmpeg"):
|
if not self.ffmpeg_ok:
|
return subprocess.CompletedProcess(command, 1, "", "synthetic failure")
|
Path(command[-1]).write_bytes(b"fLaC synthetic")
|
return subprocess.CompletedProcess(command, 0, "", "")
|
raise AssertionError(f"unexpected executable: {command[0]}")
|
|
|
class TranscribeMediaTests(unittest.TestCase):
|
def setUp(self) -> None:
|
self.executable_patch = patch(
|
"transcribe_media.shutil.which", side_effect=lambda name: f"C:/tools/{name}.exe"
|
)
|
self.executable_patch.start()
|
|
def tearDown(self) -> None:
|
self.executable_patch.stop()
|
|
def test_srt_timestamp_rounding(self) -> None:
|
self.assertEqual(transcribe_media._format_srt_timestamp(3661.2346), "01:01:01,235")
|
self.assertEqual(transcribe_media._format_srt_timestamp(-1), "00:00:00,000")
|
|
def test_end_to_end_writes_four_outputs_without_touching_source(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / "会议录像.mp4"
|
source_bytes = b"synthetic video bytes"
|
video.write_bytes(source_bytes)
|
runner = RecordingRunner()
|
model = FakeModel()
|
runtime = transcribe_media.RuntimeInfo(1, ("float16", "int8_float16"))
|
|
outputs, actual_runtime = transcribe_media.transcribe_video(
|
video,
|
language="zh",
|
runner=runner,
|
model_loader=lambda: (model, runtime),
|
)
|
|
self.assertEqual(video.read_bytes(), source_bytes)
|
self.assertEqual(actual_runtime, runtime)
|
self.assertEqual(outputs.directory, root / "会议录像.transcript")
|
for output in outputs.files:
|
self.assertTrue(output.is_file(), output)
|
self.assertIn("[00:00:03] 大家好。", outputs.text.read_text(encoding="utf-8"))
|
self.assertIn("00:00:03,120 --> 00:00:08,640", outputs.srt.read_text(encoding="utf-8"))
|
payload = json.loads(outputs.json.read_text(encoding="utf-8"))
|
self.assertEqual(payload["source"], video.name)
|
self.assertEqual(payload["device"], "cuda")
|
self.assertEqual(payload["compute_type"], "float16")
|
self.assertTrue(payload["vad_filter"])
|
self.assertEqual(len(payload["segments"]), 2)
|
self.assertEqual(model.calls[0][1], {"vad_filter": True, "language": "zh"})
|
|
ffmpeg_command = next(
|
command
|
for command in runner.commands
|
if Path(command[0]).name.lower().startswith("ffmpeg")
|
)
|
self.assertIn("-n", ffmpeg_command)
|
self.assertEqual(
|
ffmpeg_command[ffmpeg_command.index("-threads") + 1],
|
str(transcribe_media.MEDIA_DECODE_THREADS),
|
)
|
self.assertLess(ffmpeg_command.index("-threads"), ffmpeg_command.index("-i"))
|
self.assertEqual(ffmpeg_command[ffmpeg_command.index("-map") + 1], "0:a:0")
|
self.assertEqual(ffmpeg_command[ffmpeg_command.index("-ac") + 1], "1")
|
self.assertEqual(ffmpeg_command[ffmpeg_command.index("-ar") + 1], "16000")
|
self.assertEqual(ffmpeg_command[ffmpeg_command.index("-c:a") + 1], "flac")
|
|
def test_default_process_runner_uses_below_normal_creation_flags(self) -> None:
|
completed = subprocess.CompletedProcess(["tool"], 0, "", "")
|
with patch.object(transcribe_media, "_is_windows", return_value=True), patch.object(
|
transcribe_media.subprocess, "run", return_value=completed
|
) as run:
|
self.assertIs(transcribe_media._run_process(["tool"]), completed)
|
self.assertEqual(
|
run.call_args.kwargs["creationflags"],
|
getattr(subprocess, "BELOW_NORMAL_PRIORITY_CLASS", 0x00004000),
|
)
|
|
@unittest.skipUnless(os.name == "nt", "Windows named mutex and priority are required")
|
def test_named_mutex_blocks_second_process_and_restores_priority(self) -> None:
|
import ctypes
|
|
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
|
kernel32.GetCurrentProcess.restype = ctypes.c_void_p
|
kernel32.GetPriorityClass.argtypes = (ctypes.c_void_p,)
|
kernel32.GetPriorityClass.restype = ctypes.c_uint32
|
process_handle = kernel32.GetCurrentProcess()
|
original_priority = int(kernel32.GetPriorityClass(process_handle))
|
child_code = (
|
"import sys\n"
|
f"sys.path.insert(0, {str(PROJECT_DEV)!r})\n"
|
"import transcribe_media as module\n"
|
"try:\n"
|
" with module.media_host_guard():\n"
|
" pass\n"
|
"except module.MediaHostGuardError:\n"
|
" raise SystemExit(23)\n"
|
)
|
with transcribe_media.media_host_guard():
|
self.assertEqual(int(kernel32.GetPriorityClass(process_handle)), 0x00004000)
|
busy = subprocess.run([sys.executable, "-c", child_code], check=False)
|
self.assertEqual(busy.returncode, 23)
|
self.assertEqual(int(kernel32.GetPriorityClass(process_handle)), original_priority)
|
free = subprocess.run([sys.executable, "-c", child_code], check=False)
|
self.assertEqual(free.returncode, 0)
|
|
@unittest.skipUnless(os.name == "nt", "Windows named mutex is required")
|
def test_both_real_clis_reject_busy_before_path_or_external_actions(self) -> None:
|
scripts = (
|
(PROJECT_DEV / "transcribe_media.py", 1),
|
(PROJECT_DEV / "extract_ppt_slides.py", 2),
|
)
|
with transcribe_media.media_host_guard():
|
for script, expected_code in scripts:
|
with self.subTest(script=script.name):
|
result = subprocess.run(
|
[sys.executable, str(script), "definitely-missing-media.mp4"],
|
capture_output=True,
|
text=True,
|
encoding="utf-8",
|
errors="replace",
|
env={**os.environ, "PYTHONIOENCODING": "utf-8"},
|
check=False,
|
)
|
self.assertEqual(result.returncode, expected_code)
|
self.assertIn("已有重型媒体任务运行", result.stderr)
|
self.assertNotIn("视频文件不存在", result.stderr)
|
self.assertNotIn("输入视频不存在", result.stderr)
|
|
def test_win32_guard_failures_close_every_created_handle(self) -> None:
|
def kernel(**values):
|
defaults = {
|
"CreateMutexW": unittest.mock.Mock(return_value=101),
|
"WaitForSingleObject": unittest.mock.Mock(return_value=0),
|
"ReleaseMutex": unittest.mock.Mock(return_value=True),
|
"CloseHandle": unittest.mock.Mock(return_value=True),
|
"GetCurrentProcess": unittest.mock.Mock(return_value=202),
|
"GetPriorityClass": unittest.mock.Mock(return_value=0x20),
|
"SetPriorityClass": unittest.mock.Mock(return_value=True),
|
}
|
defaults.update(values)
|
return types.SimpleNamespace(**defaults)
|
|
scenarios = (
|
("wait", kernel(WaitForSingleObject=unittest.mock.Mock(return_value=0x102))),
|
("get", kernel(GetPriorityClass=unittest.mock.Mock(return_value=0))),
|
("set", kernel(SetPriorityClass=unittest.mock.Mock(return_value=False))),
|
)
|
for name, fake in scenarios:
|
with self.subTest(name=name), patch.object(
|
transcribe_media, "_is_windows", return_value=True
|
), patch.object(transcribe_media, "_windows_kernel32", return_value=fake), patch.object(
|
transcribe_media, "_win32_last_error", return_value=5
|
):
|
with self.assertRaises(transcribe_media.MediaHostGuardError):
|
with transcribe_media.media_host_guard():
|
self.fail("guard must not yield")
|
fake.CloseHandle.assert_called_once_with(101)
|
if name == "wait":
|
fake.ReleaseMutex.assert_not_called()
|
else:
|
fake.ReleaseMutex.assert_called_once_with(101)
|
|
create_failed = kernel(CreateMutexW=unittest.mock.Mock(return_value=0))
|
with patch.object(transcribe_media, "_is_windows", return_value=True), patch.object(
|
transcribe_media, "_windows_kernel32", return_value=create_failed
|
), patch.object(transcribe_media, "_win32_last_error", return_value=5):
|
with self.assertRaises(transcribe_media.MediaHostGuardError):
|
with transcribe_media.media_host_guard():
|
self.fail("guard must not yield")
|
create_failed.CloseHandle.assert_not_called()
|
|
def test_existing_output_is_not_overwritten(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / "input.mp4"
|
video.write_bytes(b"video")
|
output = root / "result"
|
output.mkdir()
|
existing = output / "input.txt"
|
existing.write_text("keep", encoding="utf-8")
|
|
with self.assertRaisesRegex(
|
transcribe_media.MediaTranscriptionError, "默认不会覆盖"
|
):
|
transcribe_media.transcribe_video(
|
video,
|
output,
|
runner=RecordingRunner(),
|
model_loader=lambda: self.fail("model should not load"),
|
)
|
self.assertEqual(existing.read_text(encoding="utf-8"), "keep")
|
|
def test_video_without_audio_fails_before_model_load(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
video = Path(temp_dir) / "silent.mp4"
|
video.write_bytes(b"video")
|
with self.assertRaisesRegex(
|
transcribe_media.MediaTranscriptionError, "不包含可用音轨"
|
):
|
transcribe_media.transcribe_video(
|
video,
|
runner=RecordingRunner(has_audio=False),
|
model_loader=lambda: self.fail("model should not load"),
|
)
|
|
def test_ffmpeg_failure_leaves_no_completed_output(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / "broken.mp4"
|
video.write_bytes(b"video")
|
with self.assertRaisesRegex(
|
transcribe_media.MediaTranscriptionError, "FFmpeg 提取第一条音轨失败"
|
):
|
transcribe_media.transcribe_video(
|
video,
|
runner=RecordingRunner(ffmpeg_ok=False),
|
model_loader=lambda: self.fail("model should not load"),
|
)
|
self.assertFalse((root / "broken.transcript").exists())
|
self.assertEqual(list(root.glob(".broken.transcribe-*")), [])
|
|
def test_four_hour_chunk_plan_has_expected_overlap_and_edges(self) -> None:
|
chunks = transcribe_media._plan_audio_chunks(14_400.0)
|
|
self.assertEqual(len(chunks), 12)
|
self.assertEqual(
|
chunks[0],
|
transcribe_media.AudioChunk(1, 12, 0.0, 1200.0, 0.0, 1205.0),
|
)
|
self.assertEqual(
|
chunks[1],
|
transcribe_media.AudioChunk(2, 12, 1200.0, 2400.0, 1195.0, 2405.0),
|
)
|
self.assertEqual(
|
chunks[-1],
|
transcribe_media.AudioChunk(
|
12, 12, 13200.0, 14400.0, 13195.0, 14400.0
|
),
|
)
|
self.assertEqual(transcribe_media._plan_audio_chunks(1200.0), ())
|
|
def test_chunk_midpoint_has_one_owner_and_global_timestamps(self) -> None:
|
first, second = transcribe_media._plan_audio_chunks(2400.0 + 1.0)[:2]
|
boundary = transcribe_media.TranscriptSegment(1198.0, 1202.0, "边界")
|
first_result = transcribe_media._project_chunk_segments(
|
(boundary,), first, 2401.0
|
)
|
second_local = transcribe_media.TranscriptSegment(3.0, 7.0, "边界")
|
second_result = transcribe_media._project_chunk_segments(
|
(second_local,), second, 2401.0
|
)
|
|
self.assertEqual(first_result, ())
|
self.assertEqual(
|
second_result,
|
(transcribe_media.TranscriptSegment(1198.0, 1202.0, "边界"),),
|
)
|
|
def test_long_audio_reuses_one_model_and_first_detected_language(self) -> None:
|
responses = [
|
[FakeSegment(10.0, 12.0, "第一块"), FakeSegment(1198.0, 1202.0, "丢弃")],
|
[FakeSegment(3.0, 7.0, "边界归第二块")],
|
[FakeSegment(4.0, 6.0, "末块")],
|
]
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / "long.mp4"
|
source_bytes = b"long synthetic video"
|
video.write_bytes(source_bytes)
|
runner = RecordingRunner(audio_duration=2401.0)
|
model = ScriptedModel(responses)
|
runtime = transcribe_media.RuntimeInfo(1, ("float16",))
|
load_count = 0
|
|
def load_model():
|
nonlocal load_count
|
load_count += 1
|
return model, runtime
|
|
outputs, _ = transcribe_media.transcribe_video(
|
video,
|
runner=runner,
|
model_loader=load_model,
|
)
|
|
self.assertEqual(load_count, 1)
|
self.assertEqual(len(model.calls), 3)
|
self.assertEqual(model.calls[0][1], {"vad_filter": True})
|
self.assertEqual(
|
model.calls[1][1], {"vad_filter": True, "language": "zh"}
|
)
|
self.assertEqual(
|
model.calls[2][1], {"vad_filter": True, "language": "zh"}
|
)
|
payload = json.loads(outputs.json.read_text(encoding="utf-8"))
|
self.assertEqual(payload["language"], "zh")
|
self.assertEqual(
|
[(item["start"], item["end"], item["text"]) for item in payload["segments"]],
|
[
|
(10.0, 12.0, "第一块"),
|
(1198.0, 1202.0, "边界归第二块"),
|
(2399.0, 2401.0, "末块"),
|
],
|
)
|
self.assertEqual(video.read_bytes(), source_bytes)
|
ffmpeg_commands = [
|
command
|
for command in runner.commands
|
if Path(command[0]).name.lower().startswith("ffmpeg")
|
]
|
self.assertEqual(len(ffmpeg_commands), 4)
|
self.assertEqual(list(root.glob(".long.transcribe-*")), [])
|
|
def test_long_audio_failure_cleans_chunk_staging_and_outputs(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / "long-fail.mp4"
|
video.write_bytes(b"video")
|
runner = RecordingRunner(audio_duration=2401.0)
|
model = ScriptedModel(
|
[[FakeSegment(1, 2, "ok")], [FakeSegment(1, 2, "never")]],
|
fail_on_call=2,
|
)
|
|
with self.assertRaisesRegex(
|
transcribe_media.MediaTranscriptionError, "GPU 语音转写失败"
|
):
|
transcribe_media.transcribe_video(
|
video,
|
runner=runner,
|
model_loader=lambda: (
|
model,
|
transcribe_media.RuntimeInfo(1, ("float16",)),
|
),
|
)
|
|
self.assertFalse((root / "long-fail.transcript").exists())
|
self.assertEqual(list(root.glob(".long-fail.transcribe-*")), [])
|
self.assertEqual(len(model.calls), 2)
|
|
def test_commit_keyboard_interrupt_rolls_back_after_each_partial_move(self) -> None:
|
for interrupt_after in (1, 2, 3):
|
with self.subTest(interrupt_after=interrupt_after):
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / f"interrupt-{interrupt_after}.mp4"
|
source_bytes = b"video"
|
video.write_bytes(source_bytes)
|
output = root / "result"
|
runner = RecordingRunner()
|
model = FakeModel()
|
original_replace = type(video).replace
|
interrupt = KeyboardInterrupt(f"after move {interrupt_after}")
|
move_count = 0
|
|
def replace_then_interrupt(path, destination):
|
nonlocal move_count
|
result = original_replace(path, destination)
|
move_count += 1
|
if move_count == interrupt_after:
|
raise interrupt
|
return result
|
|
with patch.object(
|
type(video), "replace", new=replace_then_interrupt
|
):
|
with self.assertRaises(KeyboardInterrupt) as caught:
|
transcribe_media.transcribe_video(
|
video,
|
output,
|
runner=runner,
|
model_loader=lambda: (
|
model,
|
transcribe_media.RuntimeInfo(1, ("float16",)),
|
),
|
)
|
|
self.assertIs(caught.exception, interrupt)
|
self.assertFalse(output.exists())
|
self.assertEqual(
|
list(root.glob(f".{video.stem}.transcribe-*")), []
|
)
|
self.assertEqual(video.read_bytes(), source_bytes)
|
|
def test_commit_oserror_is_wrapped_and_partial_output_is_rolled_back(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
video = root / "ordinary-error.mp4"
|
video.write_bytes(b"video")
|
output = root / "result"
|
original_replace = type(video).replace
|
move_count = 0
|
|
def fail_before_second_move(path, destination):
|
nonlocal move_count
|
move_count += 1
|
if move_count == 2:
|
raise OSError("synthetic move failure")
|
return original_replace(path, destination)
|
|
with patch.object(type(video), "replace", new=fail_before_second_move):
|
with self.assertRaisesRegex(
|
transcribe_media.MediaTranscriptionError,
|
"提交输出文件失败,已回滚已写结果",
|
):
|
transcribe_media.transcribe_video(
|
video,
|
output,
|
runner=RecordingRunner(),
|
model_loader=lambda: (
|
FakeModel(),
|
transcribe_media.RuntimeInfo(1, ("float16",)),
|
),
|
)
|
|
self.assertFalse(output.exists())
|
self.assertEqual(list(root.glob(".ordinary-error.transcribe-*")), [])
|
|
def test_cuda_model_loader_never_requests_cpu(self) -> None:
|
calls: list[tuple[str, dict[str, object]]] = []
|
|
class WhisperModel:
|
def __init__(self, name: str, **kwargs: object) -> None:
|
calls.append((name, kwargs))
|
|
fake_ctranslate2 = types.ModuleType("ctranslate2")
|
fake_ctranslate2.get_cuda_device_count = lambda: 1
|
fake_ctranslate2.get_supported_compute_types = lambda device: {"float16", "int8"}
|
fake_faster_whisper = types.ModuleType("faster_whisper")
|
fake_faster_whisper.WhisperModel = WhisperModel
|
|
with patch.dict(
|
sys.modules,
|
{
|
"ctranslate2": fake_ctranslate2,
|
"faster_whisper": fake_faster_whisper,
|
},
|
):
|
_, runtime = transcribe_media._load_cuda_model()
|
|
self.assertEqual(runtime.cuda_device_count, 1)
|
self.assertEqual(calls, [("large-v3", {"device": "cuda", "compute_type": "float16"})])
|
|
|
if __name__ == "__main__":
|
unittest.main()
|