Skip to content

Commit 8247a2a

Browse files
committed
Handle audio data URL MIME variants
1 parent 2686e2e commit 8247a2a

4 files changed

Lines changed: 40 additions & 5 deletions

File tree

src/mistral_common/protocol/instruct/chunk.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222

2323
def _strip_audio_data_url_prefix(data: str) -> str:
2424
r"""Remove the optional base64 audio data URL prefix."""
25-
if re.match(r"^data:audio/\w+;base64,", data):
25+
if re.match(r"^data:audio/[^;,]+(?:;[^;,]*)*;base64,", data):
2626
return data.split(",", 1)[1]
2727
return data
2828

src/mistral_common/tokens/tokenizers/audio.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22
import io
33
import logging
44
import math
5-
import re
65
import warnings
76
from dataclasses import dataclass
87
from enum import Enum
@@ -19,7 +18,12 @@
1918
is_soundfile_installed,
2019
is_soxr_installed,
2120
)
22-
from mistral_common.protocol.instruct.chunk import AudioChunk, AudioURLChunk, AudioURLType
21+
from mistral_common.protocol.instruct.chunk import (
22+
AudioChunk,
23+
AudioURLChunk,
24+
AudioURLType,
25+
_strip_audio_data_url_prefix,
26+
)
2327

2428
if is_soxr_installed():
2529
import soxr
@@ -114,8 +118,7 @@ def from_base64(audio_base64: str, strict: bool = True) -> "Audio":
114118
"""
115119
assert_soundfile_installed()
116120

117-
if re.match(r"^data:audio/\w+;base64,", audio_base64):
118-
audio_base64 = audio_base64.split(",")[1]
121+
audio_base64 = _strip_audio_data_url_prefix(audio_base64)
119122

120123
try:
121124
audio_bytes = base64.b64decode(audio_base64)

tests/test_audio.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -144,6 +144,15 @@ def rmse(a: np.ndarray, b: np.ndarray) -> float:
144144
raise ValueError(f"Unknown format {format}")
145145

146146

147+
@pytest.mark.parametrize("mime_type", ["audio/x-wav", "audio/vnd.wave", "audio/wav;codec=pcm"])
148+
def test_audio_from_audio_chunk_accepts_data_url_mime_variants(mime_type: str) -> None:
149+
b64 = _make_dummy_base64()
150+
audio = Audio.from_audio_chunk(AudioChunk(input_audio=f"data:{mime_type};base64,{b64}"))
151+
152+
assert audio.sampling_rate == 16000
153+
assert audio.format == "wav"
154+
155+
147156
@pytest.mark.parametrize(
148157
"freq, expected_mel",
149158
[

tests/test_converters.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import base64
22
import copy
33
import io
4+
import time
45
import warnings
56
from pathlib import Path
67
from typing import Any
@@ -44,6 +45,7 @@
4445
ImageURLChunk,
4546
TextChunk,
4647
ThinkChunk,
48+
_strip_audio_data_url_prefix,
4749
)
4850
from mistral_common.protocol.instruct.messages import (
4951
AssistantMessage,
@@ -167,6 +169,27 @@ def test_convert_input_audio_chunk() -> None:
167169
assert AudioChunk.from_openai(typeddict_openai) == chunk
168170

169171

172+
@pytest.mark.parametrize("mime_type", ["audio/x-wav", "audio/vnd.wave", "audio/wav;codec=pcm"])
173+
def test_convert_input_audio_chunk_with_data_url_mime_variants(mime_type: str) -> None:
174+
input_audio = DUMMY_AUDIO_CHUNK.input_audio
175+
assert isinstance(input_audio, str)
176+
chunk = AudioChunk(input_audio=f"data:{mime_type};base64,{input_audio}")
177+
178+
openai_dict = chunk.to_openai()
179+
180+
assert openai_dict["input_audio"]["data"] == input_audio
181+
assert openai_dict["input_audio"]["format"] == "wav"
182+
183+
184+
def test_audio_data_url_prefix_malformed_parameters_stays_linear() -> None:
185+
malformed = f"data:audio/wav{';' * 30}x"
186+
187+
start = time.monotonic()
188+
189+
assert _strip_audio_data_url_prefix(malformed) == malformed
190+
assert time.monotonic() - start < 1
191+
192+
170193
@pytest.mark.parametrize(
171194
["openai_image_url_chunk", "image_url_chunk"],
172195
[

0 commit comments

Comments
 (0)