mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
fix(proxy): match /v1/audio/speech content-type to the returned audio format
This commit is contained in:
parent
e0ed0a4c7a
commit
4e2574ce08
5 changed files with 104 additions and 11 deletions
|
|
@ -7,7 +7,12 @@ 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,
|
||||
get_file_mime_type_from_extension,
|
||||
)
|
||||
from litellm.types.utils import FileTypes
|
||||
|
||||
|
||||
|
|
@ -323,3 +328,26 @@ 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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -10877,15 +10878,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),
|
||||
|
|
|
|||
|
|
@ -0,0 +1,30 @@
|
|||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type
|
||||
|
||||
|
||||
@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_resolve_speech_media_type(upstream_content_type, response_format, expected):
|
||||
resolved = resolve_speech_media_type(
|
||||
upstream_content_type=upstream_content_type,
|
||||
response_format=response_format,
|
||||
)
|
||||
assert resolved == expected
|
||||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue