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:
2026-03-21 23:18:39 +03:00
co-authored by Claude Opus 4.6
parent 8632960354
commit 4c5399c2fc
9 changed files with 747 additions and 610 deletions
+103 -94
View File
@@ -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"),
+262 -293
View File
@@ -1,9 +1,7 @@
from collections.abc import Generator
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from huggingface_hub.errors import LocalEntryNotFoundError
from local_transcriber.transcriber import (
Segment,
@@ -15,24 +13,53 @@ from local_transcriber.transcriber import (
)
def _make_raw_segments(count: int) -> list:
"""Create mock raw segments as returned by faster-whisper."""
segments = []
for i in range(count):
seg = MagicMock()
seg.start = float(i * 5)
seg.end = float(i * 5 + 4)
seg.text = f" Segment {i}"
segments.append(seg)
return segments
# === Helpers ===
def _make_info(language: str = "ru", probability: float = 0.95, duration: float = 60.0):
info = MagicMock()
info.language = language
info.language_probability = probability
info.duration = duration
return info
def _make_result(
count: int = 2,
language: str = "ru",
probability: float = 0.95,
duration: float = 60.0,
device_used: str = "cpu",
) -> TranscribeResult:
segments = [
Segment(start=float(i * 5), end=float(i * 5 + 4), text=f" Segment {i}")
for i in range(count)
]
return TranscribeResult(
segments=segments,
language=language,
language_probability=probability,
duration=duration,
device_used=device_used,
)
def _make_backend(
model=None,
transcribe_result=None,
create_model_error=None,
transcribe_error=None,
model_path="/mock/model",
):
"""Создаёт mock-бэкенд с настраиваемым поведением."""
backend = MagicMock()
backend.ensure_model_available.return_value = model_path
if create_model_error:
backend.create_model.side_effect = create_model_error
else:
backend.create_model.return_value = model or MagicMock()
if transcribe_error:
backend.transcribe.side_effect = transcribe_error
elif transcribe_result:
backend.transcribe.return_value = transcribe_result
else:
backend.transcribe.return_value = _make_result()
return backend
def _create_model_dir(path: Path) -> Path:
@@ -45,14 +72,14 @@ def _create_model_dir(path: Path) -> Path:
return path
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_collects_segments(mock_model_cls):
raw_segments = _make_raw_segments(3)
info = _make_info()
# === transcribe() tests ===
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_collects_segments(mock_get_backend):
result_data = _make_result(count=3)
backend = _make_backend(transcribe_result=result_data)
mock_get_backend.return_value = backend
result = transcribe(
file_path=Path("test.mp3"),
@@ -68,14 +95,11 @@ def test_transcribe_collects_segments(mock_model_cls):
assert result.duration == 60.0
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_calls_on_segment(mock_model_cls):
raw_segments = _make_raw_segments(3)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_calls_on_segment(mock_get_backend):
result_data = _make_result(count=3)
backend = _make_backend(transcribe_result=result_data)
mock_get_backend.return_value = backend
callback = MagicMock()
@@ -86,28 +110,24 @@ def test_transcribe_calls_on_segment(mock_model_cls):
on_segment=callback,
)
assert callback.call_count == 3
# Each call should receive a Segment instance
for call_args in callback.call_args_list:
seg = call_args[0][0]
assert isinstance(seg, Segment)
# on_segment is passed through to backend.transcribe
call_args = backend.transcribe.call_args
assert call_args.kwargs.get("on_segment") is callback or call_args[0][3] is callback
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_cuda_fallback(mock_model_cls):
raw_segments = _make_raw_segments(2)
info = _make_info()
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_cuda_fallback(mock_get_backend):
"""CUDA error at init -> fallback на CPU."""
cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
)
# First call (cuda) raises, second call (cpu) succeeds
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
def backend_for_device(device):
return cuda_backend if device == "cuda" else cpu_backend
def model_side_effect(model_name, device, compute_type):
if device == "cuda":
raise RuntimeError("CUDA out of memory")
return cpu_instance
mock_model_cls.side_effect = model_side_effect
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
@@ -120,14 +140,12 @@ def test_transcribe_cuda_fallback(mock_model_cls):
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_device_used(mock_model_cls):
raw_segments = _make_raw_segments(1)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_device_used(mock_get_backend):
backend = _make_backend(
transcribe_result=_make_result(count=1, device_used="cuda"),
)
mock_get_backend.return_value = backend
result = transcribe(
file_path=Path("test.mp3"),
@@ -136,98 +154,47 @@ def test_transcribe_device_used(mock_model_cls):
)
assert result.device_used == "cuda"
mock_model_cls.assert_called_once_with("tiny", device="cuda", compute_type="int8")
backend.create_model.assert_called_once()
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_cuda_fallback_on_transcribe_call(mock_model_cls):
"""CUDA error in model.transcribe() (not __init__) triggers CPU fallback."""
raw_segments = _make_raw_segments(2)
info = _make_info()
cuda_instance = MagicMock()
cuda_instance.transcribe.side_effect = RuntimeError("CUDA error during transcription")
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
call_count = 0
def model_side_effect(model_name, device, compute_type):
nonlocal call_count
call_count += 1
if device == "cuda":
return cuda_instance
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_midstream_fallback_no_duplicate_callbacks(mock_model_cls):
"""on_segment is not called for partial GPU segments on mid-stream fallback."""
info = _make_info()
# GPU iterator: yields 1 segment then raises CUDA error
def _gpu_generator():
seg = MagicMock()
seg.start = 0.0
seg.end = 4.0
seg.text = " GPU seg"
yield seg
raise RuntimeError("CUDA out of memory mid-stream")
cuda_instance = MagicMock()
cuda_instance.transcribe.return_value = (_gpu_generator(), info)
cpu_segments = _make_raw_segments(2)
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(cpu_segments), info)
def model_side_effect(model_name, device, compute_type):
if device == "cuda":
return cuda_instance
return cpu_instance
mock_model_cls.side_effect = model_side_effect
callback = MagicMock()
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
on_segment=callback,
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
# callback: 1 from partial GPU pass + 2 from full CPU pass = 3
# The GPU partial segment is NOT in the final result (segments list reset),
# but on_segment was called live as segments streamed.
# This is acceptable — on_segment is a live progress callback.
# The important thing is that result.segments contains only CPU segments.
assert all(s.text.startswith(" Segment") for s in result.segments)
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_reports_missing_socksio_for_proxy(mock_model_cls):
mock_model_cls.side_effect = ImportError(
"Using SOCKS proxy, but the 'socksio' package is not installed."
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_cuda_fallback_on_transcribe_call(mock_get_backend):
"""CUDA error in transcribe (not init) triggers CPU fallback."""
cuda_backend = _make_backend(
transcribe_error=RuntimeError("CUDA error during transcription"),
)
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
)
with pytest.raises(RuntimeError, match="socksio"):
def backend_for_device(device):
return cuda_backend if device == "cuda" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_reports_missing_socksio_for_proxy(mock_get_backend):
backend = _make_backend(
create_model_error=ImportError(
"Using SOCKS proxy, but the 'socksio' package is not installed."
),
)
mock_get_backend.return_value = backend
# ImportError is not caught as backend error → propagates
with pytest.raises(ImportError, match="socksio"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
@@ -235,14 +202,10 @@ def test_transcribe_reports_missing_socksio_for_proxy(mock_model_cls):
)
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_reports_status_transitions(mock_model_cls):
raw_segments = _make_raw_segments(1)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
mock_model_cls.return_value = instance
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_reports_status_transitions(mock_get_backend):
backend = _make_backend(transcribe_result=_make_result(count=1))
mock_get_backend.return_value = backend
statuses: list[str] = []
@@ -253,14 +216,138 @@ def test_transcribe_reports_status_transitions(mock_model_cls):
on_status=statuses.append,
)
assert statuses == [
"Инициализирую модель на cpu...",
"Транскрибирую...",
"Транскрибирую... 00:04 / 01:00 [1 сегм.]",
]
# load_model reports init status, _transcribe_file reports transcribe status
assert any("Инициализирую модель" in s for s in statuses)
assert any("Транскрибирую" in s for s in statuses)
@patch("local_transcriber.transcriber.snapshot_download")
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_strict_cuda_error(mock_get_backend):
"""strict_device=True + CUDA error -> raise, без fallback."""
backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="CUDA out of memory"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_non_strict_cuda_fallback(mock_get_backend):
"""strict_device=False + CUDA error -> fallback на CPU."""
cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_backend = _make_backend(
transcribe_result=_make_result(count=2, device_used="cpu"),
model_path="/mock/cpu/model",
)
def backend_for_device(device):
return cuda_backend if device == "cuda" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=False,
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.get_backend")
def test_transcribe_strict_cuda_error_during_transcription(mock_get_backend):
"""strict_device=True + CUDA error during transcription -> raise."""
backend = _make_backend(
transcribe_error=RuntimeError("CUDA error during transcription"),
)
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="CUDA error during transcription"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
# === load_model() tests ===
@patch("local_transcriber.transcriber.get_backend")
def test_load_model_cuda_fallback(mock_get_backend):
cuda_backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
cpu_model = MagicMock()
cpu_backend = _make_backend(model=cpu_model, model_path="/mock/cpu/model")
def backend_for_device(device):
return cuda_backend if device == "cuda" else cpu_backend
mock_get_backend.side_effect = backend_for_device
with pytest.warns(UserWarning, match="Переключение на CPU"):
model, actual_device, backend, model_path = load_model("tiny", "cuda", "int8")
assert actual_device == "cpu"
assert model is cpu_model
@patch("local_transcriber.transcriber.get_backend")
def test_load_model_strict_raises(mock_get_backend):
backend = _make_backend(create_model_error=RuntimeError("CUDA out of memory"))
mock_get_backend.return_value = backend
with pytest.raises(RuntimeError, match="CUDA out of memory"):
load_model("tiny", "cuda", "int8", strict_device=True)
@patch("local_transcriber.transcriber.get_backend")
def test_load_model_returns_backend_and_path(mock_get_backend):
backend = _make_backend(model_path="/mock/model/path")
mock_get_backend.return_value = backend
model, actual_device, returned_backend, model_path = load_model("tiny", "cpu", "int8")
assert returned_backend is backend
assert model_path == "/mock/model/path"
assert actual_device == "cpu"
# === _transcribe_file() tests ===
def test__transcribe_file_basic():
result_data = _make_result(count=2)
backend = _make_backend(transcribe_result=result_data)
tfr = _transcribe_file(
model=MagicMock(),
actual_device="cpu",
backend=backend,
model_path="/mock/model",
file_path=Path("test.mp3"),
model_name="tiny",
compute_type="int8",
)
assert len(tfr.result.segments) == 2
assert tfr.actual_device == "cpu"
assert tfr.backend is backend
assert tfr.model_path == "/mock/model"
# === ensure_model_available() tests (через FasterWhisperBackend) ===
@patch("local_transcriber.backends.faster_whisper.snapshot_download")
def test_ensure_model_available_uses_cache_first(mock_snapshot_download, tmp_path):
model_dir = _create_model_dir(tmp_path / "cache-model")
mock_snapshot_download.return_value = str(model_dir)
@@ -268,22 +355,15 @@ def test_ensure_model_available_uses_cache_first(mock_snapshot_download, tmp_pat
result = ensure_model_available("large-v3")
assert result == str(model_dir)
mock_snapshot_download.assert_called_once_with(
"Systran/faster-whisper-large-v3",
local_files_only=True,
allow_patterns=[
"config.json",
"preprocessor_config.json",
"model.bin",
"tokenizer.json",
"vocabulary.*",
],
)
mock_snapshot_download.assert_called_once()
assert mock_snapshot_download.call_args.kwargs["local_files_only"] is True
@patch("local_transcriber.transcriber._validate_model_dir")
@patch("local_transcriber.transcriber.snapshot_download")
@patch("local_transcriber.backends.faster_whisper._validate_model_dir")
@patch("local_transcriber.backends.faster_whisper.snapshot_download")
def test_ensure_model_available_downloads_on_cache_miss(mock_snapshot_download, mock_validate_model_dir):
from huggingface_hub.errors import LocalEntryNotFoundError
mock_snapshot_download.side_effect = [
LocalEntryNotFoundError("not cached"),
"/downloaded/model",
@@ -295,10 +375,8 @@ def test_ensure_model_available_downloads_on_cache_miss(mock_snapshot_download,
assert result == "/downloaded/model"
assert mock_snapshot_download.call_args_list[0].kwargs["local_files_only"] is True
assert mock_snapshot_download.call_args_list[1].kwargs["local_files_only"] is False
assert statuses == [
"Проверяю кэш модели large-v3...",
"Скачиваю модель large-v3 из Hugging Face...",
]
assert "Проверяю кэш модели large-v3..." in statuses
assert "Скачиваю модель large-v3 из Hugging Face..." in statuses
def test_ensure_model_available_accepts_local_directory(tmp_path):
@@ -311,7 +389,10 @@ def test_ensure_model_available_accepts_local_directory(tmp_path):
def test_ensure_model_available_accepts_repo_id(tmp_path):
model_dir = _create_model_dir(tmp_path / "repo-model")
with patch("local_transcriber.transcriber.snapshot_download", return_value=str(model_dir)) as mock_snapshot_download:
with patch(
"local_transcriber.backends.faster_whisper.snapshot_download",
return_value=str(model_dir),
) as mock_snapshot_download:
result = ensure_model_available("org/model")
assert result == str(model_dir)
@@ -323,7 +404,7 @@ def test_ensure_model_available_rejects_unsupported_alias():
ensure_model_available("distil-large-v3")
@patch("local_transcriber.transcriber.snapshot_download")
@patch("local_transcriber.backends.faster_whisper.snapshot_download")
def test_ensure_model_available_redownloads_incomplete_cache(mock_snapshot_download, tmp_path):
incomplete = tmp_path / "incomplete"
incomplete.mkdir()
@@ -349,11 +430,7 @@ def test_ensure_model_available_redownloads_incomplete_cache(mock_snapshot_downl
result = ensure_model_available("large-v3", on_status=statuses.append)
assert result == str(complete)
assert statuses == [
"Проверяю кэш модели large-v3...",
"Кэш модели large-v3 неполный, докачиваю...",
"Скачиваю модель large-v3 из Hugging Face...",
]
assert "Кэш модели large-v3 неполный, докачиваю..." in statuses
def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path):
@@ -363,111 +440,3 @@ def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path):
with pytest.raises(ValueError, match="Неполная локальная модель"):
ensure_model_available(str(model_dir))
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_strict_cuda_error(mock_model_cls):
"""strict_device=True + CUDA error -> raise, без fallback."""
mock_model_cls.side_effect = RuntimeError("CUDA out of memory")
with pytest.raises(RuntimeError, match="CUDA out of memory"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_non_strict_cuda_fallback(mock_model_cls):
"""strict_device=False + CUDA error -> fallback на CPU."""
raw_segments = _make_raw_segments(2)
info = _make_info()
cpu_instance = MagicMock()
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
def model_side_effect(model_name, device, compute_type):
if device == "cuda":
raise RuntimeError("CUDA out of memory")
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"):
result = transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=False,
)
assert result.device_used == "cpu"
assert len(result.segments) == 2
@patch("local_transcriber.transcriber.WhisperModel")
def test_transcribe_strict_cuda_error_during_transcription(mock_model_cls):
"""strict_device=True + CUDA error during transcription -> raise."""
cuda_instance = MagicMock()
cuda_instance.transcribe.side_effect = RuntimeError("CUDA error during transcription")
mock_model_cls.return_value = cuda_instance
with pytest.raises(RuntimeError, match="CUDA error during transcription"):
transcribe(
file_path=Path("test.mp3"),
model_name="tiny",
device="cuda",
strict_device=True,
)
# === load_model tests ===
@patch("local_transcriber.transcriber.WhisperModel")
def test_load_model_cuda_fallback(mock_model_cls):
cpu_instance = MagicMock()
def model_side_effect(model_name, device, compute_type):
if device == "cuda":
raise RuntimeError("CUDA out of memory")
return cpu_instance
mock_model_cls.side_effect = model_side_effect
with pytest.warns(UserWarning, match="Переключение на CPU"):
model, actual_device = load_model("tiny", "cuda", "int8")
assert actual_device == "cpu"
assert model is cpu_instance
@patch("local_transcriber.transcriber.WhisperModel")
def test_load_model_strict_raises(mock_model_cls):
mock_model_cls.side_effect = RuntimeError("CUDA out of memory")
with pytest.raises(RuntimeError, match="CUDA out of memory"):
load_model("tiny", "cuda", "int8", strict_device=True)
@patch("local_transcriber.transcriber.WhisperModel")
def test__transcribe_file_basic(mock_model_cls):
raw_segments = _make_raw_segments(2)
info = _make_info()
instance = MagicMock()
instance.transcribe.return_value = (iter(raw_segments), info)
tfr = _transcribe_file(
model=instance,
actual_device="cpu",
file_path=Path("test.mp3"),
model_name="tiny",
compute_type="int8",
)
assert len(tfr.result.segments) == 2
assert tfr.actual_device == "cpu"
assert tfr.model is instance