diff --git a/pyrit/converter/audio_speed_converter.py b/pyrit/converter/audio_speed_converter.py index f1a6acdb11..abc5a47f40 100644 --- a/pyrit/converter/audio_speed_converter.py +++ b/pyrit/converter/audio_speed_converter.py @@ -3,6 +3,7 @@ import io import logging +import math from typing import Any, Literal import numpy as np @@ -45,13 +46,13 @@ def __init__( output_format (str): The format of the audio file, defaults to "wav". speed_factor (float): The factor by which to change the speed. Values > 1.0 speed up the audio, values < 1.0 slow it down. - Must be greater than 0 and at most 100. Defaults to 1.5. + Must be finite, greater than 0, and at most 100. Defaults to 1.5. Raises: - ValueError: If speed_factor is not positive or exceeds 100. + ValueError: If speed_factor is non-finite, not positive, or exceeds 100. """ - if speed_factor <= 0 or speed_factor > 100: - raise ValueError("speed_factor must be greater than 0 and at most 100.") + if not math.isfinite(speed_factor) or speed_factor <= 0 or speed_factor > 100: + raise ValueError("speed_factor must be finite, greater than 0, and at most 100.") self._output_format = output_format self._speed_factor = speed_factor diff --git a/tests/unit/converter/test_audio_speed_converter.py b/tests/unit/converter/test_audio_speed_converter.py index b76c920d34..a1fce82583 100644 --- a/tests/unit/converter/test_audio_speed_converter.py +++ b/tests/unit/converter/test_audio_speed_converter.py @@ -117,16 +117,23 @@ async def test_convert_async_file_not_found(): def test_invalid_speed_factor_zero(): """speed_factor of 0 should raise ValueError.""" - with pytest.raises(ValueError, match="speed_factor must be greater than 0"): + with pytest.raises(ValueError, match="speed_factor must be finite, greater than 0"): AudioSpeedConverter(speed_factor=0) def test_invalid_speed_factor_negative(): """Negative speed_factor should raise ValueError.""" - with pytest.raises(ValueError, match="speed_factor must be greater than 0"): + with pytest.raises(ValueError, match="speed_factor must be finite, greater than 0"): AudioSpeedConverter(speed_factor=-1.0) +@pytest.mark.parametrize("speed_factor", [float("nan"), float("inf"), float("-inf")]) +def test_invalid_speed_factor_non_finite(speed_factor): + """Non-finite speed factors should fail during converter construction.""" + with pytest.raises(ValueError, match="speed_factor must be finite"): + AudioSpeedConverter(speed_factor=speed_factor) + + async def test_unsupported_input_type(sqlite_instance): """Passing an unsupported input_type should raise ValueError.""" converter = AudioSpeedConverter(speed_factor=1.5)