diff --git a/pyrit/converter/audio_frequency_converter.py b/pyrit/converter/audio_frequency_converter.py index 9e851986d6..68f083e5d6 100644 --- a/pyrit/converter/audio_frequency_converter.py +++ b/pyrit/converter/audio_frequency_converter.py @@ -88,11 +88,16 @@ async def convert_async(self, *, prompt: str, input_type: PromptDataType = "audi phase = np.exp(1j * 2 * np.pi * self._shift_value * np.arange(len(data)) / sample_rate) if data.ndim > 1: phase = phase[:, np.newaxis] - shifted_data = data * phase - - # Floating-point WAV samples already use a normalized amplitude scale. - output_dtype = data.dtype if np.issubdtype(data.dtype, np.floating) else np.dtype(np.int16) - output_data = shifted_data.real.astype(output_dtype) + # 8-bit PCM is unsigned and centred on 128, so shift around that midpoint. + midpoint = 128 if data.dtype == np.uint8 else 0 + shifted_data = (data.astype(np.float64) - midpoint) * phase + + # Keep the input sample format so integer audio isn't rescaled or wrapped. + output_data = shifted_data.real + midpoint + if np.issubdtype(data.dtype, np.integer): + info = np.iinfo(data.dtype) + output_data = np.clip(output_data, info.min, info.max) + output_data = output_data.astype(data.dtype) # Write to a fresh buffer so a shorter output cannot retain input bytes. output_bytes_io = io.BytesIO() diff --git a/tests/unit/converter/test_audio_frequency_converter.py b/tests/unit/converter/test_audio_frequency_converter.py index e192422645..23aaf68b43 100644 --- a/tests/unit/converter/test_audio_frequency_converter.py +++ b/tests/unit/converter/test_audio_frequency_converter.py @@ -38,6 +38,46 @@ async def test_frequency_preserves_float_waveform_async( np.testing.assert_allclose(output, expected, atol=1e-7) +@pytest.mark.usefixtures("sqlite_instance") +@pytest.mark.parametrize( + ("dtype", "samples", "expected_shifted"), + [ + (np.int16, [1 << 14, 1 << 13, -(1 << 14), -(1 << 13)], [1 << 14, 0, 1 << 14, 0]), + (np.int32, [1 << 30, 1 << 29, -(1 << 30), -(1 << 29)], [1 << 30, 0, 1 << 30, 0]), + (np.uint8, [192, 160, 64, 96], [192, 128, 192, 128]), + ], +) +@pytest.mark.parametrize("stereo", [False, True]) +@pytest.mark.parametrize("shift_value", [0, 2000]) +async def test_frequency_preserves_integer_waveform_async( + *, + tmp_path: Path, + dtype: type[np.integer], + samples: list[int], + expected_shifted: list[int], + stereo: bool, + shift_value: int, +) -> None: + """Integer WAVs keep their sample format and midpoint instead of being cast to int16.""" + midpoint = 128 if dtype is np.uint8 else 0 + data = np.array(samples, dtype=dtype) + expected = data if shift_value == 0 else np.array(expected_shifted, dtype=dtype) + if stereo: + # Mirror the second channel around the format's midpoint. + data = np.column_stack((data, (2 * midpoint - data.astype(np.int64)).astype(dtype))) + expected = np.column_stack((expected, (2 * midpoint - expected.astype(np.int64)).astype(dtype))) + source = tmp_path / "int.wav" + await asyncio.to_thread(wavfile.write, source, 8000, data) + + result = await AudioFrequencyConverter(shift_value=shift_value).convert_async(prompt=str(source)) + rate, output = await asyncio.to_thread(wavfile.read, result.output_text) + + assert rate == 8000 + assert output.dtype == data.dtype + assert output.shape == data.shape + np.testing.assert_allclose(output, expected, atol=1) + + async def test_convert_async_success(sqlite_instance): # Simulate WAV data sample_rate = 44100