From 4c5399c2fca2646f21c85ceaf088328b492783f4 Mon Sep 17 00:00:00 2001 From: Dmitry Dementev Date: Sat, 21 Mar 2026 23:18:39 +0300 Subject: [PATCH] =?UTF-8?q?refactor(transcriber):=20=D0=B2=D0=B2=D0=B5?= =?UTF-8?q?=D0=B4=D0=B5=D0=BD=D0=B0=20pluggable-=D0=B0=D1=80=D1=85=D0=B8?= =?UTF-8?q?=D1=82=D0=B5=D0=BA=D1=82=D1=83=D1=80=D0=B0=20=D0=B1=D1=8D=D0=BA?= =?UTF-8?q?=D0=B5=D0=BD=D0=B4=D0=BE=D0=B2=20=D1=82=D1=80=D0=B0=D0=BD=D1=81?= =?UTF-8?q?=D0=BA=D1=80=D0=B8=D0=BF=D1=86=D0=B8=D0=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Зачем: - подготовка к добавлению 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) --- src/local_transcriber/backends/__init__.py | 28 + src/local_transcriber/backends/base.py | 46 ++ .../backends/faster_whisper.py | 181 ++++++ src/local_transcriber/cli.py | 61 +- src/local_transcriber/formatter.py | 2 +- src/local_transcriber/transcriber.py | 253 ++------ src/local_transcriber/types.py | 34 ++ tests/test_cli.py | 197 ++++--- tests/test_transcriber.py | 555 +++++++++--------- 9 files changed, 747 insertions(+), 610 deletions(-) create mode 100644 src/local_transcriber/backends/__init__.py create mode 100644 src/local_transcriber/backends/base.py create mode 100644 src/local_transcriber/backends/faster_whisper.py create mode 100644 src/local_transcriber/types.py diff --git a/src/local_transcriber/backends/__init__.py b/src/local_transcriber/backends/__init__.py new file mode 100644 index 0000000..213475c --- /dev/null +++ b/src/local_transcriber/backends/__init__.py @@ -0,0 +1,28 @@ +"""Реестр бэкендов транскрипции и выбор бэкенда по устройству.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from .base import Backend + + +def get_backend(device: str) -> Backend: + """Возвращает экземпляр бэкенда для указанного устройства. + + Импорты ленивые — бэкенд загружается только при запросе. + """ + if device == "openvino": + try: + from .openvino import OpenVINOBackend + except ImportError: + raise ValueError( + "OpenVINO бэкенд недоступен. Установите: pip install openvino-genai" + ) from None + return OpenVINOBackend() + + # cuda, cpu и всё остальное → faster-whisper + from .faster_whisper import FasterWhisperBackend + + return FasterWhisperBackend() diff --git a/src/local_transcriber/backends/base.py b/src/local_transcriber/backends/base.py new file mode 100644 index 0000000..0435064 --- /dev/null +++ b/src/local_transcriber/backends/base.py @@ -0,0 +1,46 @@ +"""Протокол бэкенда транскрипции.""" + +from __future__ import annotations + +from collections.abc import Callable +from pathlib import Path +from typing import Any, Protocol + +from local_transcriber.types import Segment, TranscribeResult + + +class Backend(Protocol): + """Минимальный интерфейс бэкенда транскрипции. + + Бэкенды реализуют этот протокол (structural typing) — + наследование не требуется. + """ + + def ensure_model_available( + self, + model_name: str, + compute_type: str, + on_status: Callable[[str], None] | None = None, + ) -> str: + """Гарантирует наличие модели, возвращает путь к файлам.""" + ... + + def create_model( + self, + model_path: str, + device: str, + compute_type: str, + ) -> Any: + """Создаёт модель. Возвращает backend-специфичный объект.""" + ... + + def transcribe( + self, + model: Any, + file_path: Path, + language: str | None, + on_segment: Callable[[Segment], None] | None = None, + on_status: Callable[[str], None] | None = None, + ) -> TranscribeResult: + """Транскрибирует файл, возвращает результат.""" + ... diff --git a/src/local_transcriber/backends/faster_whisper.py b/src/local_transcriber/backends/faster_whisper.py new file mode 100644 index 0000000..8998e54 --- /dev/null +++ b/src/local_transcriber/backends/faster_whisper.py @@ -0,0 +1,181 @@ +"""Бэкенд транскрипции на основе faster-whisper (CTranslate2).""" + +from __future__ import annotations + +import gc +import io +import warnings +from collections.abc import Callable +from pathlib import Path +from typing import Any + +# CUDA bootstrap — должен быть ДО импорта faster_whisper / ctranslate2 +from local_transcriber._cuda_bootstrap import ensure_cublas_loadable + +ensure_cublas_loadable() + +from faster_whisper import WhisperModel # noqa: E402 +from huggingface_hub import snapshot_download # noqa: E402 +from huggingface_hub.errors import LocalEntryNotFoundError # noqa: E402 + +from local_transcriber.types import Segment, TranscribeResult # noqa: E402 + +MODEL_REPOS = { + "tiny": "Systran/faster-whisper-tiny", + "base": "Systran/faster-whisper-base", + "small": "Systran/faster-whisper-small", + "medium": "Systran/faster-whisper-medium", + "large-v3": "Systran/faster-whisper-large-v3", +} + +MODEL_ALLOW_PATTERNS = [ + "config.json", + "preprocessor_config.json", + "model.bin", + "tokenizer.json", + "vocabulary.*", +] + +MODEL_REQUIRED_FILES = [ + "config.json", + "model.bin", + "tokenizer.json", +] + + +class FasterWhisperBackend: + """Бэкенд транскрипции через faster-whisper (CTranslate2).""" + + def ensure_model_available( + self, + model_name: str, + compute_type: str, + on_status: Callable[[str], None] | None = None, + ) -> str: + """Резолвит alias модели в repo_id и гарантирует наличие файлов.""" + local_path = Path(model_name).expanduser() + if local_path.is_dir(): + _validate_model_dir(local_path) + return str(local_path) + + repo_id = _resolve_model_repo(model_name) + + try: + _notify(on_status, f"Проверяю кэш модели {model_name}...") + cached_path = Path(_snapshot_download(repo_id, local_files_only=True)) + _validate_model_dir(cached_path) + return str(cached_path) + except LocalEntryNotFoundError: + pass + except ValueError: + _notify(on_status, f"Кэш модели {model_name} неполный, докачиваю...") + + _notify(on_status, f"Скачиваю модель {model_name} из Hugging Face...") + downloaded_path = Path(_snapshot_download(repo_id, local_files_only=False)) + _validate_model_dir(downloaded_path) + return str(downloaded_path) + + def create_model( + self, + model_path: str, + device: str, + compute_type: str, + ) -> Any: + """Создаёт WhisperModel.""" + try: + return WhisperModel(model_path, device=device, compute_type=compute_type) + except ImportError as exc: + if _is_missing_socksio_error(exc): + raise RuntimeError( + "Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, " + "нужная для загрузки модели из Hugging Face через proxy. " + "Обновите окружение: `uv sync`." + ) from exc + raise + + def transcribe( + self, + model: Any, + file_path: Path, + language: str | None, + on_segment: Callable[[Segment], None] | None = None, + on_status: Callable[[str], None] | None = None, + ) -> TranscribeResult: + """Транскрибирует файл через faster-whisper.""" + segment_generator, info = model.transcribe( + str(file_path), language=language, + ) + total_duration = info.duration + segments: list[Segment] = [] + for raw_seg in segment_generator: + seg = Segment(start=raw_seg.start, end=raw_seg.end, text=raw_seg.text) + if on_segment is not None: + on_segment(seg) + segments.append(seg) + _notify( + on_status, + f"Транскрибирую... {_fmt_time(seg.end)} / {_fmt_time(total_duration)}" + f" [{len(segments)} сегм.]", + ) + + return TranscribeResult( + segments=segments, + language=info.language, + language_probability=info.language_probability, + duration=info.duration, + device_used="", # оркестратор проставит actual_device + ) + + +def _notify(on_status: Callable[[str], None] | None, message: str) -> None: + if on_status is not None: + on_status(message) + + +def _fmt_time(seconds: float) -> str: + m, s = divmod(int(seconds), 60) + h, m = divmod(m, 60) + return f"{h}:{m:02d}:{s:02d}" if h else f"{m:02d}:{s:02d}" + + +def _resolve_model_repo(model_name: str) -> str: + if "/" in model_name: + return model_name + repo_id = MODEL_REPOS.get(model_name) + if repo_id is None: + expected = ", ".join(MODEL_REPOS) + raise ValueError(f"Неподдерживаемая модель '{model_name}'. Ожидалось одно из: {expected}") + return repo_id + + +def _snapshot_download(repo_id: str, local_files_only: bool) -> str: + try: + return snapshot_download( + repo_id, + local_files_only=local_files_only, + allow_patterns=MODEL_ALLOW_PATTERNS, + ) + except ImportError as exc: + if _is_missing_socksio_error(exc): + raise RuntimeError( + "Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, " + "нужная для загрузки модели из Hugging Face через proxy. " + "Обновите окружение: `uv sync`." + ) from exc + raise + + +def _validate_model_dir(model_dir: Path) -> None: + missing = [ + filename for filename in MODEL_REQUIRED_FILES if not (model_dir / filename).exists() + ] + if not any(model_dir.glob("vocabulary.*")): + missing.append("vocabulary.*") + if missing: + missing_str = ", ".join(missing) + raise ValueError(f"Неполная локальная модель в '{model_dir}': отсутствуют {missing_str}") + + +def _is_missing_socksio_error(exc: BaseException) -> bool: + msg = str(exc).lower() + return "socks proxy" in msg and "socksio" in msg diff --git a/src/local_transcriber/cli.py b/src/local_transcriber/cli.py index dd45a2c..a62215e 100644 --- a/src/local_transcriber/cli.py +++ b/src/local_transcriber/cli.py @@ -14,9 +14,7 @@ from .transcriber import ( Segment, _is_cuda_error, _transcribe_file, - ensure_model_available, load_model, - transcribe, ) from .utils import ( build_output_path, @@ -31,6 +29,16 @@ app = typer.Typer() console = Console(stderr=True) +def _format_device_info(device_used: str) -> str: + """Формирует строку устройства для шапки транскрипта.""" + if device_used == "cuda": + gpu_name = get_gpu_name() + return f"CUDA ({gpu_name or 'Unknown GPU'})" + if device_used == "openvino": + return "OpenVINO (CPU)" + return "CPU" + + @app.command() def main( files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"), @@ -42,7 +50,8 @@ def main( ), output: Path | None = typer.Option(None, "--output", "-o", help="Путь к выходному файлу"), device: str | None = typer.Option( - None, "--device", "-d", show_default=False, help="Устройство (auto|cpu|cuda) [по умолч.: auto]" + None, "--device", "-d", show_default=False, + help="Устройство (auto|cpu|cuda) [по умолч.: auto]" ), compute_type: str | None = typer.Option( None, "--compute-type", show_default=False, @@ -120,7 +129,6 @@ def _run_single( validated_file = validate_input_file(file) requested_device = defaults["device"] resolved_device = detect_device(requested_device) - # Если пользователь явно указал устройство — запрещаем fallback на CPU strict = requested_device != "auto" output_path = build_output_path(validated_file, output) @@ -131,15 +139,11 @@ def _run_single( f"Compute: [bold]{defaults['compute_type']}[/bold]" ) - model_path = ensure_model_available( - defaults["model"], on_status=lambda message: console.print(message) - ) - def on_segment(seg: Segment) -> None: console.print(f" [{seg.start:.2f}s] {seg.text.strip()}") - model_obj, actual_device = load_model( - model_path, resolved_device, defaults["compute_type"], + model_obj, actual_device, backend, model_path = load_model( + defaults["model"], resolved_device, defaults["compute_type"], on_status=lambda msg: console.print(msg), strict_device=strict, ) @@ -147,8 +151,10 @@ def _run_single( tfr = _transcribe_file( model=model_obj, actual_device=actual_device, + backend=backend, + model_path=model_path, file_path=validated_file, - model_name=model_path, + model_name=defaults["model"], compute_type=defaults["compute_type"], language=defaults["language"] if defaults["language"] != "auto" else None, on_segment=on_segment if verbose else None, @@ -176,12 +182,7 @@ def _run_single( f"Речь не обнаружена в файле {validated_file.name}", style="yellow" ) - if result.device_used == "cuda": - gpu_name = get_gpu_name() - device_info = f"CUDA ({gpu_name or 'Unknown GPU'})" - else: - device_info = "CPU" - + device_info = _format_device_info(result.device_used) language_mode = "detected" if defaults["language"] == "auto" else "forced" content = format_transcript( @@ -232,15 +233,12 @@ def _run_batch( raise SystemExit(1) return - # Phase 2: Load model + # Phase 2: Load model (ensure + create в одном вызове) requested_device = defaults["device"] resolved_device = detect_device(requested_device) strict = requested_device != "auto" - model_path = ensure_model_available( - defaults["model"], on_status=lambda msg: console.print(msg) - ) - model_obj, actual_device = load_model( - model_path, resolved_device, defaults["compute_type"], + model_obj, actual_device, backend, model_path = load_model( + defaults["model"], resolved_device, defaults["compute_type"], on_status=lambda msg: console.print(msg), strict_device=strict, ) @@ -277,8 +275,10 @@ def _run_batch( tfr = _transcribe_file( model=model_obj, actual_device=actual_device, + backend=backend, + model_path=model_path, file_path=file, - model_name=model_path, + model_name=defaults["model"], compute_type=defaults["compute_type"], language=defaults["language"] if defaults["language"] != "auto" else None, on_segment=on_segment if verbose else None, @@ -291,8 +291,11 @@ def _run_batch( f" {file.name}: fallback на {tfr.actual_device} при транскрипции", style="yellow", ) - # Обновляем после возможного mid-stream fallback на CPU - model_obj, actual_device = tfr.model, tfr.actual_device + # Обновляем после возможного mid-stream fallback + model_obj = tfr.model + actual_device = tfr.actual_device + backend = tfr.backend + model_path = tfr.model_path result = tfr.result @@ -301,11 +304,7 @@ def _run_batch( f" Речь не обнаружена: {file.name}", style="yellow" ) - if result.device_used == "cuda": - gpu_name = get_gpu_name() - device_info = f"CUDA ({gpu_name or 'Unknown GPU'})" - else: - device_info = "CPU" + device_info = _format_device_info(result.device_used) content = format_transcript( result=result, diff --git a/src/local_transcriber/formatter.py b/src/local_transcriber/formatter.py index 609ac08..0233d51 100644 --- a/src/local_transcriber/formatter.py +++ b/src/local_transcriber/formatter.py @@ -4,7 +4,7 @@ from dataclasses import dataclass from datetime import datetime from pathlib import Path -from .transcriber import Segment, TranscribeResult +from .types import Segment, TranscribeResult _PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы _MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца diff --git a/src/local_transcriber/transcriber.py b/src/local_transcriber/transcriber.py index 8fd72ec..cc43a66 100644 --- a/src/local_transcriber/transcriber.py +++ b/src/local_transcriber/transcriber.py @@ -1,65 +1,18 @@ -"""Обёртка над faster-whisper: загрузка моделей, транскрипция, CUDA fallback.""" +"""Оркестрация транскрипции: выбор бэкенда, загрузка модели, fallback.""" import warnings from collections.abc import Callable -from dataclasses import dataclass from pathlib import Path +from typing import Any -# Должен быть ДО импорта faster_whisper / ctranslate2 -from local_transcriber._cuda_bootstrap import ensure_cublas_loadable +from local_transcriber.backends import get_backend -ensure_cublas_loadable() - -from faster_whisper import WhisperModel # noqa: E402 -from huggingface_hub import snapshot_download -from huggingface_hub.errors import LocalEntryNotFoundError - -MODEL_REPOS = { - "tiny": "Systran/faster-whisper-tiny", - "base": "Systran/faster-whisper-base", - "small": "Systran/faster-whisper-small", - "medium": "Systran/faster-whisper-medium", - "large-v3": "Systran/faster-whisper-large-v3", -} - -# allow — фильтр для snapshot_download (какие файлы скачивать из репозитория); -# required — для валидации (что обязано быть после скачивания/в локальной модели) -MODEL_ALLOW_PATTERNS = [ - "config.json", - "preprocessor_config.json", - "model.bin", - "tokenizer.json", - "vocabulary.*", -] - -MODEL_REQUIRED_FILES = [ - "config.json", - "model.bin", - "tokenizer.json", -] - - -@dataclass -class Segment: - start: float # seconds - end: float # seconds - text: str - - -@dataclass -class TranscribeResult: - segments: list[Segment] - language: str - language_probability: float - duration: float # seconds - device_used: str # "cpu" / "cuda" - - -@dataclass -class TranscribeFileResult: - result: TranscribeResult - model: WhisperModel - actual_device: str +# Re-export из types.py для обратной совместимости +from local_transcriber.types import ( # noqa: F401 + Segment, + TranscribeFileResult, + TranscribeResult, +) def load_model( @@ -68,15 +21,21 @@ def load_model( compute_type: str, on_status: Callable[[str], None] | None = None, strict_device: bool = False, -) -> tuple[WhisperModel, str]: - """Загружает модель с CUDA-фолбеком. Возвращает (model, actual_device).""" +) -> tuple[Any, str, Any, str]: + """Загружает модель: ensure + create с fallback. + + Возвращает (model, actual_device, backend, model_path). + """ + backend = get_backend(device) actual_device = device + + model_path = backend.ensure_model_available(model_name, compute_type, on_status) + try: _notify_status(on_status, f"Инициализирую модель на {device}...") - model = _create_model(model_name, device, compute_type) + model = backend.create_model(model_path, device, compute_type) except (RuntimeError, ValueError) as exc: - # strict — пользователь явно указал устройство, fallback запрещён - if device != "cpu" and _is_cuda_error(exc): + if device != "cpu" and _is_backend_error(exc, device): if strict_device: raise warnings.warn( @@ -85,16 +44,21 @@ def load_model( stacklevel=2, ) actual_device = "cpu" + backend = get_backend("cpu") + model_path = backend.ensure_model_available(model_name, compute_type, on_status) _notify_status(on_status, "Инициализирую модель на cpu...") - model = _create_model(model_name, "cpu", compute_type) + model = backend.create_model(model_path, "cpu", compute_type) else: raise - return model, actual_device + + return model, actual_device, backend, model_path def _transcribe_file( - model: WhisperModel, + model: Any, actual_device: str, + backend: Any, + model_path: str, file_path: Path, model_name: str, compute_type: str, @@ -103,39 +67,40 @@ def _transcribe_file( on_status: Callable[[str], None] | None = None, strict_device: bool = False, ) -> TranscribeFileResult: - """Транскрибирует один файл. При mid-stream CUDA fallback перезагружает модель.""" + """Транскрибирует один файл. При mid-stream fallback перезагружает модель.""" lang_arg = language if language and language != "auto" else None try: _notify_status(on_status, "Транскрибирую...") - segments, info = _run_transcription(model, file_path, lang_arg, on_segment, on_status) + result = backend.transcribe(model, file_path, lang_arg, on_segment, on_status) + result.device_used = actual_device except (RuntimeError, ValueError) as exc: - # Mid-stream fallback: GPU может упасть с OOM уже во время транскрипции, - # поэтому перезагружаем модель на CPU и начинаем сначала - if actual_device != "cpu" and _is_cuda_error(exc): + if actual_device != "cpu" and _is_backend_error(exc, actual_device): if strict_device: raise warnings.warn( - f"CUDA ошибка при транскрипции: {exc}. " + f"Ошибка при транскрипции на {actual_device}: {exc}. " "Переключение на CPU и повтор.", stacklevel=2, ) actual_device = "cpu" + backend = get_backend("cpu") + model_path = backend.ensure_model_available(model_name, compute_type, on_status) _notify_status(on_status, "Инициализирую модель на cpu...") - model = _create_model(model_name, "cpu", compute_type) + model = backend.create_model(model_path, "cpu", compute_type) _notify_status(on_status, "Транскрибирую...") - segments, info = _run_transcription(model, file_path, lang_arg, on_segment, on_status) + result = backend.transcribe(model, file_path, lang_arg, on_segment, on_status) + result.device_used = actual_device else: raise - result = TranscribeResult( - segments=segments, - language=info.language, - language_probability=info.language_probability, - duration=info.duration, - device_used=actual_device, + return TranscribeFileResult( + result=result, + model=model, + actual_device=actual_device, + backend=backend, + model_path=model_path, ) - return TranscribeFileResult(result=result, model=model, actual_device=actual_device) def transcribe( @@ -149,9 +114,12 @@ def transcribe( strict_device: bool = False, ) -> TranscribeResult: """High-level API: загрузка модели + транскрипция за один вызов.""" - model, actual_device = load_model(model_name, device, compute_type, on_status, strict_device) + model, actual_device, backend, model_path = load_model( + model_name, device, compute_type, on_status, strict_device, + ) tfr = _transcribe_file( - model, actual_device, file_path, model_name, compute_type, + model, actual_device, backend, model_path, + file_path, model_name, compute_type, language, on_segment, on_status, strict_device, ) return tfr.result @@ -159,127 +127,30 @@ def transcribe( def ensure_model_available( model_name: str, + device: str = "cpu", + compute_type: str = "float32", on_status: Callable[[str], None] | None = None, ) -> str: - """Резолвит alias модели в repo_id и гарантирует наличие файлов. - - Стратегия: cache-first (``local_files_only=True``), затем download. - Два вызова ``snapshot_download`` — чтобы не лезть в сеть, если модель уже в кэше. - """ - local_path = Path(model_name).expanduser() - if local_path.is_dir(): - _validate_model_dir(local_path) - return str(local_path) - - repo_id = _resolve_model_repo(model_name) - - try: - _notify_status(on_status, f"Проверяю кэш модели {model_name}...") - cached_path = Path(_snapshot_download(repo_id, local_files_only=True)) - _validate_model_dir(cached_path) - return str(cached_path) - except LocalEntryNotFoundError: - pass - except ValueError: - _notify_status(on_status, f"Кэш модели {model_name} неполный, докачиваю...") - - _notify_status(on_status, f"Скачиваю модель {model_name} из Hugging Face...") - downloaded_path = Path(_snapshot_download(repo_id, local_files_only=False)) - _validate_model_dir(downloaded_path) - return str(downloaded_path) - - -def _run_transcription(model, file_path, lang_arg, on_segment, on_status=None): - """Run model.transcribe and iterate segments. Returns (segments, info).""" - segment_generator, info = model.transcribe(str(file_path), language=lang_arg) - total_duration = info.duration - segments: list[Segment] = [] - for raw_seg in segment_generator: - seg = Segment(start=raw_seg.start, end=raw_seg.end, text=raw_seg.text) - if on_segment is not None: - on_segment(seg) - segments.append(seg) - _notify_status( - on_status, - f"Транскрибирую... {_fmt_time(seg.end)} / {_fmt_time(total_duration)}" - f" [{len(segments)} сегм.]", - ) - return segments, info - - -def _create_model(model_name: str, device: str, compute_type: str): - try: - return WhisperModel(model_name, device=device, compute_type=compute_type) - except ImportError as exc: - # WhisperModel при инициализации может загружать файлы через HF Hub; - # если в системе настроен SOCKS proxy, но socksio не установлен, - # HF Hub бросает ImportError — оборачиваем в понятное сообщение - if _is_missing_socksio_error(exc): - raise RuntimeError( - "Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, " - "нужная для загрузки модели из Hugging Face через proxy. " - "Обновите окружение: `uv sync`." - ) from exc - raise + """Публичный helper: гарантирует наличие модели для указанного бэкенда.""" + backend = get_backend(device) + return backend.ensure_model_available(model_name, compute_type, on_status) def _is_cuda_error(exc: BaseException) -> bool: + """Проверка CUDA ошибок — используется в cli.py для Windows-диагностики.""" msg = str(exc).lower() return any(k in msg for k in ("cuda", "cublas", "cudnn", "out of memory")) -def _is_missing_socksio_error(exc: BaseException) -> bool: - msg = str(exc).lower() - return "socks proxy" in msg and "socksio" in msg - - -def _fmt_time(seconds: float) -> str: - m, s = divmod(int(seconds), 60) - h, m = divmod(m, 60) - return f"{h}:{m:02d}:{s:02d}" if h else f"{m:02d}:{s:02d}" +def _is_backend_error(exc: BaseException, device: str) -> bool: + """Определяет, связана ли ошибка с конкретным бэкендом (а не с пользовательскими данными).""" + if device in ("cuda", "cpu"): + return _is_cuda_error(exc) + # openvino и другие бэкенды: конкретные паттерны ошибок добавим + # при реализации бэкенда; пока — не маскируем ошибки + return False def _notify_status(on_status: Callable[[str], None] | None, message: str) -> None: if on_status is not None: on_status(message) - - -def _resolve_model_repo(model_name: str) -> str: - if "/" in model_name: - return model_name - - repo_id = MODEL_REPOS.get(model_name) - if repo_id is None: - expected = ", ".join(MODEL_REPOS) - raise ValueError(f"Неподдерживаемая модель '{model_name}'. Ожидалось одно из: {expected}") - - return repo_id - - -def _snapshot_download(repo_id: str, local_files_only: bool) -> str: - try: - return snapshot_download( - repo_id, - local_files_only=local_files_only, - allow_patterns=MODEL_ALLOW_PATTERNS, - ) - except ImportError as exc: - if _is_missing_socksio_error(exc): - raise RuntimeError( - "Обнаружен SOCKS proxy, но не установлена зависимость `socksio`, " - "нужная для загрузки модели из Hugging Face через proxy. " - "Обновите окружение: `uv sync`." - ) from exc - raise - - -def _validate_model_dir(model_dir: Path) -> None: - missing = [ - filename for filename in MODEL_REQUIRED_FILES if not (model_dir / filename).exists() - ] - if not any(model_dir.glob("vocabulary.*")): - missing.append("vocabulary.*") - - if missing: - missing_str = ", ".join(missing) - raise ValueError(f"Неполная локальная модель в '{model_dir}': отсутствуют {missing_str}") diff --git a/src/local_transcriber/types.py b/src/local_transcriber/types.py new file mode 100644 index 0000000..a8eca66 --- /dev/null +++ b/src/local_transcriber/types.py @@ -0,0 +1,34 @@ +"""Общие типы данных для всех бэкендов транскрипции.""" + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + + +@dataclass +class Segment: + start: float # seconds + end: float # seconds + text: str + + +@dataclass +class TranscribeResult: + segments: list[Segment] + language: str + language_probability: float + duration: float # seconds + device_used: str # "cpu" / "cuda" / "openvino" + + +@dataclass +class TranscribeFileResult: + result: TranscribeResult + model: Any # backend-specific model handle + actual_device: str + backend: Any = None # backend instance (для переиспользования в батче) + model_path: str = "" # путь к модели (меняется при cross-backend fallback) + + +StatusCallback = Callable[[str], None] | None +SegmentCallback = Callable[[Segment], None] | None diff --git a/tests/test_cli.py b/tests/test_cli.py index 2cef38b..ff6c7d0 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -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"), diff --git a/tests/test_transcriber.py b/tests/test_transcriber.py index 6775f48..b39a9bd 100644 --- a/tests/test_transcriber.py +++ b/tests/test_transcriber.py @@ -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