diff --git a/tests/test_cli.py b/tests/test_cli.py new file mode 100644 index 0000000..d1f1afc --- /dev/null +++ b/tests/test_cli.py @@ -0,0 +1,143 @@ +"""Command-line interface: output naming and error reporting.""" + +import importlib +import shutil +import sys + +import numpy as np +import pytest + +from whisper_mlx.audio import AudioLoadError, load_audio + +cli = importlib.import_module("whisper_mlx.cli") + + +@pytest.fixture +def run_cli(monkeypatch, tmp_path): + """Run `vayu ` with transcription and writing stubbed out. + + Returns (exit_code, names written, inputs passed to transcribe()). + """ + + def run(*args, transcribe=None): + written, transcribed = [], [] + + def fake_transcribe(audio, **kwargs): + transcribed.append(audio) + return {"text": "", "segments": [], "language": "en"} + + monkeypatch.setattr(cli, "transcribe", transcribe or fake_transcribe) + monkeypatch.setattr( + cli, + "get_writer", + lambda fmt, out: lambda r, name, **k: written.append(name), + ) + monkeypatch.setattr(sys, "argv", ["vayu", *args, "-o", str(tmp_path)]) + try: + cli.main() + code = 0 + except SystemExit as e: + code = e.code + return code, written, transcribed + + return run + + +def test_each_input_gets_its_own_output_name(run_cli): + code, written, _ = run_cli("talks/a.mp3", "b.wav") + + assert code == 0 + assert written == ["a", "b"] + + +def test_output_name_applies_to_a_single_input(run_cli): + code, written, _ = run_cli("a.mp3", "--output-name", "custom") + + assert (code, written) == (0, ["custom"]) + + +def test_output_name_with_several_inputs_is_rejected(run_cli, capsys): + code, written, _ = run_cli("a.mp3", "b.mp3", "--output-name", "custom") + + assert code == 2 + assert written == [] + assert "--output-name" in capsys.readouterr().err + + +def test_inputs_with_the_same_stem_are_rejected(run_cli, capsys): + code, written, _ = run_cli("day1/talk.mp3", "day2/talk.mp3") + + assert code == 2 + assert written == [] + assert "talk" in capsys.readouterr().err + + +def test_stdin_is_read_inside_error_handling(run_cli, monkeypatch, capsys): + def broken_stdin(**kwargs): + raise AudioLoadError("Failed to load audio: invalid data") + + monkeypatch.setattr(cli.audio, "load_audio", broken_stdin) + + code, written, _ = run_cli("-") + + assert code == 1 + assert " - -: AudioLoadError" in capsys.readouterr().out + + +def test_stdin_output_is_named_content(run_cli, monkeypatch): + monkeypatch.setattr(cli.audio, "load_audio", lambda **k: np.zeros(16000)) + + code, written, transcribed = run_cli("-") + + assert (code, written) == (0, ["content"]) + assert isinstance(transcribed[0], np.ndarray) + + +def test_missing_file_is_reported_and_others_still_run(run_cli, capsys): + def transcribe(audio, **kwargs): + if audio == "missing.wav": + raise FileNotFoundError(f"Audio file not found: {audio}") + return {"text": "", "segments": [], "language": "en"} + + code, written, _ = run_cli("missing.wav", "present.wav", transcribe=transcribe) + + assert code == 1 + assert written == ["present"] + out = capsys.readouterr().out + assert "1 file(s) failed" in out + assert "missing.wav: FileNotFoundError" in out + + +def test_load_audio_reports_missing_file(tmp_path): + with pytest.raises(FileNotFoundError, match="Audio file not found"): + load_audio(str(tmp_path / "nope.wav")) + + +def test_load_audio_reports_missing_ffmpeg(tmp_path, monkeypatch): + clip = tmp_path / "clip.wav" + clip.write_bytes(b"not audio") + monkeypatch.setenv("PATH", str(tmp_path)) # no ffmpeg in here + + with pytest.raises(AudioLoadError, match="ffmpeg was not found"): + load_audio(str(clip)) + + +@pytest.mark.skipif(shutil.which("ffmpeg") is None, reason="needs ffmpeg") +def test_load_audio_reports_undecodable_file(tmp_path): + clip = tmp_path / "clip.mp3" + clip.write_text("this is not an mp3") + + with pytest.raises(AudioLoadError, match="Failed to load audio"): + load_audio(str(clip)) + + +def test_transcribe_checks_the_file_before_loading_the_model(monkeypatch, tmp_path): + module = importlib.import_module("whisper_mlx.transcribe") + + def no_model(*args, **kwargs): + raise AssertionError("model should not be loaded") + + monkeypatch.setattr(module.ModelHolder, "get_model", no_model) + + with pytest.raises(FileNotFoundError): + module.transcribe(str(tmp_path / "missing.mp3")) diff --git a/whisper_mlx/__init__.py b/whisper_mlx/__init__.py index 95710f1..84d7403 100644 --- a/whisper_mlx/__init__.py +++ b/whisper_mlx/__init__.py @@ -47,6 +47,7 @@ N_SAMPLES, SAMPLE_RATE, TOKENS_PER_SECOND, + AudioLoadError, load_audio, log_mel_spectrogram, pad_or_trim, @@ -91,6 +92,7 @@ "DecodingResult", # Audio processing "load_audio", + "AudioLoadError", "log_mel_spectrogram", "pad_or_trim", "SAMPLE_RATE", diff --git a/whisper_mlx/audio.py b/whisper_mlx/audio.py index 77b32bc..aea99ee 100644 --- a/whisper_mlx/audio.py +++ b/whisper_mlx/audio.py @@ -21,6 +21,10 @@ TOKENS_PER_SECOND = SAMPLE_RATE // N_SAMPLES_PER_TOKEN # 20ms per audio token +class AudioLoadError(RuntimeError): + """ffmpeg is missing or could not decode the audio.""" + + def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bool = False) -> mx.array: """ Open an audio file and read as mono waveform, resampling as necessary @@ -35,7 +39,14 @@ def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bo Returns ------- - A NumPy array containing the audio waveform, in float32 dtype. + An mx.array containing the audio waveform, in float32 dtype. + + Raises + ------ + FileNotFoundError + If the audio file does not exist. + AudioLoadError + If ffmpeg is not installed or fails to decode the audio. """ # This launches a subprocess to decode audio while down-mixing @@ -44,7 +55,7 @@ def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bo cmd = ["ffmpeg", "-i", "pipe:0"] else: if not os.path.isfile(file): - raise ValueError(f"Audio file not found: {file}") + raise FileNotFoundError(f"Audio file not found: {file}") file = os.path.realpath(file) cmd = ["ffmpeg", "-nostdin", "-i", file] @@ -60,8 +71,12 @@ def load_audio(file: Optional[str] = None, sr: int = SAMPLE_RATE, from_stdin: bo # fmt: on try: out = run(cmd, capture_output=True, check=True).stdout + except FileNotFoundError as e: + raise AudioLoadError( + "ffmpeg was not found on PATH; install it (e.g. `brew install ffmpeg`)" + ) from e except CalledProcessError as e: - raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e + raise AudioLoadError(f"Failed to load audio: {e.stderr.decode()}") from e return mx.array(np.frombuffer(out, np.int16)).flatten().astype(mx.float32) / 32768.0 diff --git a/whisper_mlx/cli.py b/whisper_mlx/cli.py index 54aa15f..c4e2c68 100644 --- a/whisper_mlx/cli.py +++ b/whisper_mlx/cli.py @@ -6,8 +6,7 @@ import pathlib import sys import warnings -from dataclasses import dataclass, field -from subprocess import CalledProcessError +from dataclasses import dataclass from typing import List from . import audio @@ -268,51 +267,62 @@ def main(): errors: List[TranscriptionError] = [] - for audio_obj in args.pop("audio"): - if audio_obj == "-": - # receive the contents from stdin rather than read a file - audio_obj = audio.load_audio(from_stdin=True) + audio_paths: List[str] = args.pop("audio") + if output_name is not None and len(audio_paths) > 1: + parser.error("--output-name can only be used with a single audio input") + # "-" reads from stdin; its output is named "content" + output_names = [ + output_name or ("content" if path == "-" else pathlib.Path(path).stem) + for path in audio_paths + ] + duplicates = sorted({n for n in output_names if output_names.count(n) > 1}) + if duplicates: + parser.error( + f"several inputs would write to the same output name: {', '.join(duplicates)}" + ) - output_name = output_name or "content" - else: - output_name = output_name or pathlib.Path(audio_obj).stem + for audio_path, name in zip(audio_paths, output_names): try: + audio_obj = ( + audio.load_audio(from_stdin=True) if audio_path == "-" else audio_path + ) result = transcribe( audio_obj, path_or_hf_repo=path_or_hf_repo, batch_size=batch_size, **args, ) - writer(result, output_name, **writer_args) + writer(result, name, **writer_args) except FileNotFoundError as e: - logger.error(f"File not found: {audio_obj}") - errors.append(TranscriptionError(audio_obj, "FileNotFoundError", str(e))) + logger.error(f"File not found: {audio_path}") + errors.append(TranscriptionError(audio_path, "FileNotFoundError", str(e))) if strict: sys.exit(1) - except CalledProcessError as e: - # FFmpeg or other subprocess failures - stderr_msg = e.stderr.decode() if e.stderr else str(e) - logger.error(f"Audio processing failed for {audio_obj}: {stderr_msg}") - errors.append(TranscriptionError(audio_obj, "CalledProcessError", stderr_msg)) + except audio.AudioLoadError as e: + # ffmpeg missing or unable to decode the input + logger.error(f"Audio processing failed for {audio_path}: {e}") + errors.append(TranscriptionError(audio_path, "AudioLoadError", str(e))) if strict: sys.exit(1) except ValueError as e: # Input validation errors - logger.error(f"Invalid input for {audio_obj}: {e}") - errors.append(TranscriptionError(audio_obj, "ValueError", str(e))) + logger.error(f"Invalid input for {audio_path}: {e}") + errors.append(TranscriptionError(audio_path, "ValueError", str(e))) if strict: sys.exit(1) - except MemoryError as e: + except MemoryError: # Out of memory - always fatal - logger.error(f"Out of memory processing {audio_obj}. Try reducing --batch-size.") + logger.error( + f"Out of memory processing {audio_path}. Try reducing --batch-size." + ) raise except KeyboardInterrupt: logger.info("Interrupted by user") sys.exit(130) except Exception as e: # Catch-all for unexpected errors - logger.exception(f"Unexpected error processing {audio_obj}") - errors.append(TranscriptionError(audio_obj, type(e).__name__, str(e))) + logger.exception(f"Unexpected error processing {audio_path}") + errors.append(TranscriptionError(audio_path, type(e).__name__, str(e))) if strict: raise diff --git a/whisper_mlx/transcribe.py b/whisper_mlx/transcribe.py index 8299ea2..945471e 100644 --- a/whisper_mlx/transcribe.py +++ b/whisper_mlx/transcribe.py @@ -1,5 +1,6 @@ # Copyright © 2023 Apple Inc. +import os import sys import warnings from typing import Any, Dict, List, Optional, Tuple, Union @@ -143,6 +144,9 @@ def transcribe( if isinstance(audio, str): if not audio: raise ValueError("Audio path cannot be empty") + if not os.path.isfile(audio): + # fail before loading (and possibly downloading) the model + raise FileNotFoundError(f"Audio file not found: {audio}") elif isinstance(audio, (np.ndarray, mx.array)): if audio.size == 0: raise ValueError("Audio array cannot be empty")