refactor(transcriber): введена pluggable-архитектура бэкендов транскрипции
- Зачем: - подготовка к добавлению OpenVINO бэкенда для ускорения на x86 CPU без CUDA. - архитектура должна позволять добавлять новые бэкенды (CoreML, AMD XDNA) без переписывания кода. - Что: - создан types.py с общими типами (Segment, TranscribeResult, TranscribeFileResult). - создан backends/base.py с Backend Protocol (3 метода: ensure_model_available, create_model, transcribe). - создан backends/faster_whisper.py — текущий код вынесен из transcriber.py в FasterWhisperBackend. - transcriber.py переделан в оркестратор: load_model() владеет полным пайплайном (ensure + create), CLI больше не вызывает ensure_model_available() отдельно. - TranscribeFileResult расширен полями backend и model_path для корректного cross-backend fallback в батч-режиме. - device_used проставляется оркестратором, а не бэкендом. - cli.py: вынесен _format_device_info(), подготовлен к openvino. - тесты обновлены: mock-точки перенесены с WhisperModel на get_backend/бэкенд-объекты. - Проверка: - uv run pytest -v — 98 passed. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
+103
-94
@@ -24,12 +24,21 @@ def _make_model():
|
||||
return MagicMock(name="WhisperModel")
|
||||
|
||||
|
||||
def _make_tfr(result=None, model=None, actual_device="cpu"):
|
||||
def _make_backend():
|
||||
return MagicMock(name="Backend")
|
||||
|
||||
|
||||
def _make_tfr(result=None, model=None, actual_device="cpu", backend=None, model_path="/models/medium"):
|
||||
if result is None:
|
||||
result = _make_result()
|
||||
if model is None:
|
||||
model = _make_model()
|
||||
return TranscribeFileResult(result=result, model=model, actual_device=actual_device)
|
||||
if backend is None:
|
||||
backend = _make_backend()
|
||||
return TranscribeFileResult(
|
||||
result=result, model=model, actual_device=actual_device,
|
||||
backend=backend, model_path=model_path,
|
||||
)
|
||||
|
||||
|
||||
def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
|
||||
@@ -37,13 +46,13 @@ def _single_patches(result=None, tmp_file=None, actual_device="cpu"):
|
||||
if result is None:
|
||||
result = _make_result(device_used=actual_device)
|
||||
model = _make_model()
|
||||
tfr = TranscribeFileResult(result=result, model=model, actual_device=actual_device)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, actual_device=actual_device, backend=backend)
|
||||
return [
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=tmp_file),
|
||||
patch("local_transcriber.cli.detect_device", return_value=actual_device),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, actual_device)),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, actual_device, backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
]
|
||||
@@ -54,7 +63,7 @@ def test_cli_happy_path_exit_code_zero(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
|
||||
patches = _single_patches(tmp_file=audio)
|
||||
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6]:
|
||||
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
|
||||
out = runner.invoke(app, [str(audio)])
|
||||
|
||||
assert out.exit_code == 0
|
||||
@@ -65,22 +74,22 @@ def test_cli_default_options_passed_to_transcribe(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio)])
|
||||
|
||||
call_kwargs = mock_transcribe_file.call_args[1]
|
||||
assert call_kwargs["model_name"] == "/models/medium"
|
||||
assert call_kwargs["model_name"] == "medium"
|
||||
assert call_kwargs["compute_type"] == "float32"
|
||||
assert call_kwargs["language"] == "ru"
|
||||
assert call_kwargs["on_segment"] is None # verbose=False
|
||||
@@ -91,15 +100,15 @@ def test_cli_custom_options(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result(device_used="cuda")
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model, actual_device="cuda")
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, actual_device="cuda", backend=backend)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/small"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/small")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
|
||||
@@ -113,7 +122,7 @@ def test_cli_custom_options(tmp_path):
|
||||
])
|
||||
|
||||
call_kwargs = mock_transcribe_file.call_args[1]
|
||||
assert call_kwargs["model_name"] == "/models/small"
|
||||
assert call_kwargs["model_name"] == "small"
|
||||
assert call_kwargs["language"] == "ru"
|
||||
assert call_kwargs["compute_type"] == "float16"
|
||||
|
||||
@@ -123,15 +132,15 @@ def test_cli_verbose_passes_on_segment_callback(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -148,7 +157,7 @@ def test_cli_empty_speech_warning(tmp_path):
|
||||
result = _make_result(segments=[])
|
||||
|
||||
patches = _single_patches(result=result, tmp_file=audio)
|
||||
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6]:
|
||||
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
|
||||
out = runner.invoke(app, [str(audio)])
|
||||
|
||||
assert out.exit_code == 0
|
||||
@@ -162,14 +171,14 @@ def test_cli_default_output_path(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript", mock_write),
|
||||
):
|
||||
@@ -187,14 +196,14 @@ def test_cli_custom_output_path(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript", mock_write),
|
||||
):
|
||||
@@ -209,15 +218,15 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -228,28 +237,27 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path):
|
||||
assert callable(call_kwargs["on_status"])
|
||||
|
||||
|
||||
def test_cli_resolves_model_before_transcribe(tmp_path):
|
||||
def test_cli_load_model_called_with_model_name(tmp_path):
|
||||
"""load_model receives model name from defaults, handles ensure internally."""
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/large-v3"))
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3") as mock_ensure,
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio), "--model", "large-v3"])
|
||||
|
||||
mock_ensure.assert_called_once()
|
||||
call_kwargs = mock_transcribe_file.call_args[1]
|
||||
assert call_kwargs["model_name"] == "/models/large-v3"
|
||||
assert mock_load_model.call_args[0][0] == "large-v3"
|
||||
|
||||
|
||||
def test_cli_windows_cuda_diagnostic(tmp_path):
|
||||
@@ -257,13 +265,13 @@ def test_cli_windows_cuda_diagnostic(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
|
||||
patch("local_transcriber.cli.sys") as mock_sys,
|
||||
):
|
||||
@@ -280,13 +288,13 @@ def test_cli_linux_cuda_error_no_windows_hint(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")),
|
||||
patch("local_transcriber.cli.sys") as mock_sys,
|
||||
):
|
||||
@@ -303,14 +311,14 @@ def test_cli_device_fallback_warning(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result(device_used="cpu")
|
||||
model = _make_model()
|
||||
tfr = TranscribeFileResult(result=result, model=model, actual_device="cpu")
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, actual_device="cpu", backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -325,15 +333,15 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path):
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result(device_used="cuda")
|
||||
model = _make_model()
|
||||
tfr = TranscribeFileResult(result=result, model=model, actual_device="cuda")
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, actual_device="cuda", backend=backend)
|
||||
mock_transcribe_file = MagicMock(return_value=tfr)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
|
||||
@@ -344,15 +352,14 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path):
|
||||
|
||||
mock_transcribe_file.reset_mock()
|
||||
result_cpu = _make_result(device_used="cpu")
|
||||
tfr_cpu = TranscribeFileResult(result=result_cpu, model=model, actual_device="cpu")
|
||||
tfr_cpu = _make_tfr(result=result_cpu, model=model, backend=backend)
|
||||
mock_transcribe_file.return_value = tfr_cpu
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", mock_transcribe_file),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -366,13 +373,13 @@ def test_cli_keyboard_interrupt(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -399,13 +406,13 @@ def test_cli_unexpected_error_verbose_traceback(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -420,13 +427,13 @@ def test_cli_unexpected_error_no_verbose_hint(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -448,14 +455,14 @@ def test_cli_batch_two_files(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -470,19 +477,18 @@ def test_cli_batch_skips_existing(tmp_path):
|
||||
b = tmp_path / "b.mp3"
|
||||
a.write_bytes(b"fake")
|
||||
b.write_bytes(b"fake")
|
||||
# Create transcript for a
|
||||
(tmp_path / "a-transcript.md").write_text("existing")
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -524,14 +530,14 @@ def test_cli_batch_force_overwrites(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -550,7 +556,8 @@ def test_cli_batch_per_file_error(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
call_count = 0
|
||||
|
||||
def transcribe_side_effect(**kwargs):
|
||||
@@ -564,8 +571,7 @@ def test_cli_batch_per_file_error(tmp_path):
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -584,7 +590,8 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
|
||||
def validate_side_effect(p):
|
||||
if not p.exists():
|
||||
@@ -595,8 +602,7 @@ def test_cli_batch_invalid_in_prescan(tmp_path):
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=validate_side_effect),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -634,42 +640,46 @@ def test_cli_config_applied(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
result = _make_result()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/tiny"))
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/tiny") as mock_ensure,
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio)])
|
||||
|
||||
mock_ensure.assert_called_once_with("tiny", on_status=mock_ensure.call_args[1]["on_status"])
|
||||
# load_model receives model name from config
|
||||
assert mock_load_model.call_args[0][0] == "tiny"
|
||||
|
||||
|
||||
def test_cli_cli_overrides_config(tmp_path):
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
model = _make_model()
|
||||
backend = _make_backend()
|
||||
result = _make_result()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/small"))
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={"model": "tiny"}),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/small") as mock_ensure,
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio), "--model", "small"])
|
||||
|
||||
mock_ensure.assert_called_once_with("small", on_status=mock_ensure.call_args[1]["on_status"])
|
||||
# CLI --model overrides config
|
||||
assert mock_load_model.call_args[0][0] == "small"
|
||||
|
||||
|
||||
def test_cli_batch_fallback_warning(tmp_path):
|
||||
@@ -681,14 +691,14 @@ def test_cli_batch_fallback_warning(tmp_path):
|
||||
|
||||
result = _make_result(device_used="cpu")
|
||||
model = _make_model()
|
||||
tfr = TranscribeFileResult(result=result, model=model, actual_device="cpu")
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, actual_device="cpu", backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -707,15 +717,15 @@ def test_cli_batch_empty_speech_warning(tmp_path):
|
||||
result_empty = _make_result(segments=[])
|
||||
result_ok = _make_result()
|
||||
model = _make_model()
|
||||
tfr_empty = _make_tfr(result=result_empty, model=model)
|
||||
tfr_ok = _make_tfr(result=result_ok, model=model)
|
||||
backend = _make_backend()
|
||||
tfr_empty = _make_tfr(result=result_empty, model=model, backend=backend)
|
||||
tfr_ok = _make_tfr(result=result_ok, model=model, backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model, "cpu", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -735,17 +745,16 @@ def test_cli_batch_midstream_fallback_warning(tmp_path):
|
||||
|
||||
model_gpu = _make_model()
|
||||
model_cpu = _make_model()
|
||||
backend = _make_backend()
|
||||
result = _make_result(device_used="cpu")
|
||||
# First file triggers mid-stream fallback
|
||||
tfr_fallback = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu")
|
||||
tfr_ok = TranscribeFileResult(result=result, model=model_cpu, actual_device="cpu")
|
||||
tfr_fallback = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
|
||||
tfr_ok = _make_tfr(result=result, model=model_cpu, actual_device="cpu", backend=backend)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda")),
|
||||
patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda", backend, "/models/medium")),
|
||||
patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
@@ -763,14 +772,14 @@ def test_cli_batch_model_loaded_once(tmp_path):
|
||||
|
||||
result = _make_result()
|
||||
model = _make_model()
|
||||
tfr = _make_tfr(result=result, model=model)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu"))
|
||||
backend = _make_backend()
|
||||
tfr = _make_tfr(result=result, model=model, backend=backend)
|
||||
mock_load_model = MagicMock(return_value=(model, "cpu", backend, "/models/medium"))
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.load_config", return_value={}),
|
||||
patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"),
|
||||
patch("local_transcriber.cli.load_model", mock_load_model),
|
||||
patch("local_transcriber.cli._transcribe_file", return_value=tfr),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
|
||||
Reference in New Issue
Block a user