Merge remote-tracking branch 'origin/litellm_fix_audio_speech_content_type' into litellm_fix_gemini_tts_container

This commit is contained in:
mateo-berri 2026-08-29 16:27:27 -07:00
commit a85e16f731
7 changed files with 234 additions and 14 deletions

View file

@ -7,7 +7,13 @@ import os
from dataclasses import dataclass
from typing import Final
from litellm.types.files import get_file_mime_type_from_extension
from litellm.types.files import (
AUDIO_FILE_TYPES,
FILE_EXTENSIONS,
FILE_MIME_TYPES,
FileType,
get_file_mime_type_from_extension,
)
from litellm.types.utils import FileTypes
@ -323,3 +329,75 @@ def calculate_request_duration(file: FileTypes) -> float | None:
except Exception:
# Silently fail if duration extraction fails
return None
DEFAULT_SPEECH_MEDIA_TYPE: Final = "audio/mpeg"
def _speech_media_type_for_response_format(response_format: str) -> str | None:
file_type: Final = next(
(candidate for candidate, extensions in FILE_EXTENSIONS.items() if response_format.lower() in extensions),
None,
)
if file_type is None or file_type not in AUDIO_FILE_TYPES:
return None
return FILE_MIME_TYPES[file_type]
def resolve_speech_media_type(upstream_content_type: str | None, response_format: str | None) -> str:
upstream_media_type: Final = (upstream_content_type or "").split(";", 1)[0].strip().lower()
if upstream_media_type.startswith("audio/"):
return upstream_media_type
requested_media_type: Final = (
None if response_format is None else _speech_media_type_for_response_format(response_format)
)
return requested_media_type or DEFAULT_SPEECH_MEDIA_TYPE
_OGG_OPUS_HEAD_WINDOW: Final = 64
_ADTS_SYNC_AND_LAYER_MASK: Final = 0xF6
_ADTS_SYNC_AND_LAYER: Final = 0xF0
_ADTS_SAMPLE_RATE_INDEX_LIMIT: Final = 13
_MPEG_SYNC_MASK: Final = 0xE0
_MPEG_LAYER_MASK: Final = 0x06
_MPEG_RESERVED_VERSION: Final = 0x01
_MPEG_INVALID_BITRATE_INDEX: Final = 0x0F
_MPEG_RESERVED_SAMPLE_RATE_INDEX: Final = 0x03
def _adts_aac_frame_media_type(header: bytes) -> str | None:
sample_rate_index: Final = (header[2] >> 2) & 0x0F
return FILE_MIME_TYPES[FileType.AAC] if sample_rate_index < _ADTS_SAMPLE_RATE_INDEX_LIMIT else None
def _mpeg_audio_frame_media_type(header: bytes) -> str | None:
version: Final = (header[1] >> 3) & 0x03
layer: Final = header[1] & _MPEG_LAYER_MASK
bitrate_index: Final = header[2] >> 4
sample_rate_index: Final = (header[2] >> 2) & 0x03
if (
(header[1] & _MPEG_SYNC_MASK) != _MPEG_SYNC_MASK
or version == _MPEG_RESERVED_VERSION
or layer == 0
or bitrate_index == _MPEG_INVALID_BITRATE_INDEX
or sample_rate_index == _MPEG_RESERVED_SAMPLE_RATE_INDEX
):
return None
return FILE_MIME_TYPES[FileType.MP3]
def speech_media_type_from_audio_bytes(audio: bytes) -> str | None:
if audio[:4] == b"RIFF" and audio[8:12] == b"WAVE":
return FILE_MIME_TYPES[FileType.WAV]
if audio[:4] == b"fLaC":
return FILE_MIME_TYPES[FileType.FLAC]
if audio[:4] == b"OggS":
is_opus: Final = b"OpusHead" in audio[:_OGG_OPUS_HEAD_WINDOW]
return FILE_MIME_TYPES[FileType.OPUS if is_opus else FileType.OGG]
if audio[:3] == b"ID3":
return FILE_MIME_TYPES[FileType.MP3]
if len(audio) < 3 or audio[0] != 0xFF:
return None
if (audio[1] & _ADTS_SYNC_AND_LAYER_MASK) == _ADTS_SYNC_AND_LAYER:
return _adts_aac_frame_media_type(audio)
return _mpeg_audio_frame_media_type(audio)

View file

@ -11,6 +11,9 @@ from typing import TYPE_CHECKING, Any, Final, Union
import httpx
from litellm.litellm_core_utils.audio_utils.utils import (
speech_media_type_from_audio_bytes,
)
from litellm.llms.base_llm.text_to_speech.transformation import (
BaseTextToSpeechConfig,
TextToSpeechRequestData,
@ -457,12 +460,11 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase):
if not response_content:
raise ValueError("No audioContent in Vertex AI TTS response")
# Decode base64 to get binary content
binary_data: Final = base64.b64decode(response_content)
# Create an httpx.Response object with the binary data
media_type: Final = speech_media_type_from_audio_bytes(binary_data)
response: Final = httpx.Response(
status_code=200,
headers={} if media_type is None else {"content-type": media_type},
content=binary_data,
)

View file

@ -263,6 +263,7 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
validated_max_agentic_loops,
)
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_litellm_metadata_from_kwargs,
@ -10975,15 +10976,11 @@ async def audio_speech(
if callback_headers:
custom_headers.update(callback_headers)
# Determine media type based on model type
media_type = "audio/mpeg" # Default for OpenAI TTS
request_model: Final = data.get("model", "")
if request_model:
request_model_lower: Final = request_model.lower()
if "gemini" in request_model_lower and (
"tts" in request_model_lower or "preview-tts" in request_model_lower
):
media_type = "audio/wav" # Gemini TTS returns WAV format after conversion
requested_format: Final = data.get("response_format")
media_type: Final = resolve_speech_media_type(
upstream_content_type=response.response.headers.get("content-type"),
response_format=requested_format if isinstance(requested_format, str) else None,
)
return StreamingResponse(
_audio_speech_chunk_generator(response),

View file

@ -347,3 +347,65 @@ class TestNormalizeTranscriptionLanguageToBcp47:
)
assert normalize_transcription_language_to_bcp47(language) == expected
class TestResolveSpeechMediaType:
@pytest.mark.parametrize(
("upstream_content_type", "response_format", "expected"),
[
("audio/wav", None, "audio/wav"),
("AUDIO/WAV", None, "audio/wav"),
("audio/flac; charset=binary", "mp3", "audio/flac"),
("application/json", "flac", "audio/flac"),
("application/octet-stream", "pcm", "audio/pcm"),
(None, "wav", "audio/wav"),
(None, "WAV", "audio/wav"),
(None, "opus", "audio/opus"),
(None, "aac", "audio/aac"),
(None, "mp3", "audio/mpeg"),
(None, "mp4", "audio/mpeg"),
(None, "bogus", "audio/mpeg"),
(None, None, "audio/mpeg"),
("", None, "audio/mpeg"),
],
)
def test_resolution(self, upstream_content_type, response_format, expected):
from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type
resolved = resolve_speech_media_type(
upstream_content_type=upstream_content_type,
response_format=response_format,
)
assert resolved == expected
class TestSpeechMediaTypeFromAudioBytes:
@pytest.mark.parametrize(
("audio", "expected"),
[
(b"RIFF\x24\x00\x00\x00WAVEfmt ", "audio/wav"),
(b"fLaC\x00\x00\x00\x22", "audio/flac"),
(b"OggS" + b"\x00" * 24 + b"OpusHead", "audio/opus"),
(b"OggS" + b"\x00" * 24 + b"\x01vorbis", "audio/ogg"),
(b"ID3\x04\x00\x00\x00\x00\x00\x00", "audio/mpeg"),
(b"\xff\xfb\x90\x64", "audio/mpeg"),
(b"\xff\xf3\x80\x00", "audio/mpeg"),
(b"\xff\xf1\x50\x80", "audio/aac"),
(b"\xff\xf9\x50\x80", "audio/aac"),
(b"RIFF\x24\x00\x00\x00AVI LIST", None),
(b"\xff\xff\xff\xff\xff\xff", None),
(b"\xff\xfb\xf0\x00", None),
(b"\xff\xfb\x9c\x00", None),
(b"\xff\xeb\x90\x00", None),
(b"\xff\xf1\xf4\x80", None),
(b"\xff\x00\x00\x00", None),
(b"\x00\x01\x02\x03\x04\x05", None),
(b"\xff\xfb", None),
(b"\xff", None),
(b"", None),
],
)
def test_sniffing(self, audio, expected):
from litellm.litellm_core_utils.audio_utils.utils import speech_media_type_from_audio_bytes
assert speech_media_type_from_audio_bytes(audio) == expected

View file

@ -1,3 +1,4 @@
import base64
from unittest.mock import MagicMock, Mock, patch
import httpx
@ -126,6 +127,48 @@ class TestVertexAITextToSpeechConfig:
assert voice_dict == voice_input
@pytest.mark.parametrize(
("audio", "expected_content_type"),
[
(b"RIFF\x24\x00\x00\x00WAVEfmt \x10\x00\x00\x00", "audio/wav"),
(b"\xff\xfb\x90\x64\x00\x00\x00\x00", "audio/mpeg"),
(b"OggS" + b"\x00" * 24 + b"OpusHead", "audio/opus"),
(b"fLaC\x00\x00\x00\x22", "audio/flac"),
],
)
def test_transform_text_to_speech_response_labels_content_type(audio, expected_content_type):
raw_response = httpx.Response(
status_code=200,
json={"audioContent": base64.b64encode(audio).decode()},
)
result = VertexAITextToSpeechConfig().transform_text_to_speech_response(
model="vertex_ai/chirp",
raw_response=raw_response,
logging_obj=MagicMock(),
)
assert result.response.headers["content-type"] == expected_content_type
assert result.response.content == audio
def test_transform_text_to_speech_response_leaves_unknown_bytes_unlabeled():
raw_pcm = b"\x00\x01\x02\x03\x04\x05\x06\x07"
raw_response = httpx.Response(
status_code=200,
json={"audioContent": base64.b64encode(raw_pcm).decode()},
)
result = VertexAITextToSpeechConfig().transform_text_to_speech_response(
model="vertex_ai/chirp",
raw_response=raw_response,
logging_obj=MagicMock(),
)
assert "content-type" not in result.response.headers
assert result.response.content == raw_pcm
@patch("litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post")
@patch.object(VertexAITextToSpeechConfig, "_ensure_access_token")
@patch.object(VertexAITextToSpeechConfig, "_get_token_and_url")

View file

@ -12,13 +12,15 @@ from __future__ import annotations
import io
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from litellm.proxy import proxy_server
@pytest.fixture
def patched_speech(monkeypatch):
def patched_speech(monkeypatch, request):
upstream_content_type = getattr(request, "param", "audio/mpeg")
monkeypatch.setattr(proxy_server, "llm_router", MagicMock())
monkeypatch.setattr(
proxy_server,
@ -37,6 +39,11 @@ def patched_speech(monkeypatch):
monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data)
class _FakeBinaryResp:
response = httpx.Response(
status_code=200,
headers={} if upstream_content_type is None else {"content-type": upstream_content_type},
)
async def aiter_bytes(self, chunk_size: int = 8192):
async def _gen():
yield b"\x00\x01\x02"
@ -152,6 +159,35 @@ def test_audio_speech_happy_path(client, auth_as, patched_speech, path):
}
@pytest.mark.parametrize(
("patched_speech", "response_format", "expected_content_type"),
[
("audio/wav", "wav", "audio/wav"),
("audio/flac", "flac", "audio/flac"),
("audio/pcm", "pcm", "audio/pcm"),
("audio/wav", "mp3", "audio/wav"),
("application/json", "flac", "audio/flac"),
(None, "wav", "audio/wav"),
(None, None, "audio/mpeg"),
],
indirect=["patched_speech"],
)
def test_audio_speech_content_type_matches_audio_format(
client, auth_as, patched_speech, response_format, expected_content_type
):
"""Regression for LIT-6482: /v1/audio/speech mislabeled wav/flac/pcm as audio/mpeg."""
payload = {
"model": "tts-1",
"input": "Hi",
"voice": "alloy",
**({} if response_format is None else {"response_format": response_format}),
}
with auth_as():
response = client.post("/v1/audio/speech", json=payload)
assert response.status_code == 200
assert response.headers.get("content-type", "").split(";")[0] == expected_content_type
@pytest.mark.parametrize("path", ["/v1/audio/speech", "/audio/speech"])
def test_audio_speech_error(client, auth_as, patched_speech_error, path):
"""Pins ``POST /v1/audio/speech`` and ``POST /audio/speech`` (error)."""

View file

@ -2,6 +2,7 @@ import asyncio
import os
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi.testclient import TestClient
@ -29,6 +30,7 @@ def _make_mock_tts_response():
inner = MagicMock()
inner.aiter_bytes = _aiter_bytes
inner._hidden_params = {}
inner.response = httpx.Response(status_code=200, headers={"content-type": "audio/mpeg"})
async def _resolver():
return inner