feat: Add audio parameter support to gemini tts models (#11287)

* feat: Add Gemini TTS audio parameter support

- Add is_model_gemini_audio_model() method to detect TTS models
- Include 'audio' parameter in supported params for TTS models
- Map OpenAI audio parameter to Gemini speechConfig format
- Add _extract_audio_response_from_parts() method to transform audio
  output to openai format

* updated unit-test to use pcm16

* - created typedict for speechconfig
- simplified gemini tts model detection
- moved gemini_tts test to test_litellm

* simplified is_model_gemini_audio_model more
This commit is contained in:
AyrennC
2025-05-31 16:20:19 -07:00
committed by GitHub
parent 13dc757873
commit 8ae79178ae
4 changed files with 349 additions and 8 deletions
+40 -1
View File
@@ -6,7 +6,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_image_obj,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import ContentType, PartType
from litellm.types.llms.vertex_ai import ContentType, PartType, SpeechConfig, VoiceConfig, PrebuiltVoiceConfig
from litellm.utils import supports_reasoning
from ...vertex_ai.gemini.transformation import _gemini_convert_messages_with_history
@@ -67,6 +67,9 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
def get_config(cls):
return super().get_config()
def is_model_gemini_audio_model(self, model: str) -> bool:
return "tts" in model
def get_supported_openai_params(self, model: str) -> List[str]:
supported_params = [
"temperature",
@@ -89,6 +92,8 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
if supports_reasoning(model):
supported_params.append("reasoning_effort")
supported_params.append("thinking")
if self.is_model_gemini_audio_model(model):
supported_params.append("audio")
return supported_params
def map_openai_params(
@@ -98,6 +103,40 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
model: str,
drop_params: bool,
) -> Dict:
# Handle audio parameter for TTS models
if self.is_model_gemini_audio_model(model):
for param, value in non_default_params.items():
if param == "audio" and isinstance(value, dict):
# Validate audio format - Gemini TTS only supports pcm16
audio_format = value.get("format")
if audio_format is not None and audio_format != "pcm16":
raise ValueError(
f"Unsupported audio format for Gemini TTS models: {audio_format}. "
f"Gemini TTS models only support 'pcm16' format as they return audio data in L16 PCM format. "
f"Please set audio format to 'pcm16'."
)
# Map OpenAI audio parameter to Gemini speech config
speech_config: SpeechConfig = {}
if "voice" in value:
prebuilt_voice_config: PrebuiltVoiceConfig = {
"voiceName": value["voice"]
}
voice_config: VoiceConfig = {
"prebuiltVoiceConfig": prebuilt_voice_config
}
speech_config["voiceConfig"] = voice_config
if speech_config:
optional_params["speechConfig"] = speech_config
# Ensure audio modality is set
if "responseModalities" not in optional_params:
optional_params["responseModalities"] = ["AUDIO"]
elif "AUDIO" not in optional_params["responseModalities"]:
optional_params["responseModalities"].append("AUDIO")
if litellm.vertex_ai_safety_settings is not None:
optional_params["safety_settings"] = litellm.vertex_ai_safety_settings
return super().map_openai_params(
@@ -2,6 +2,7 @@
## httpx client for vertex ai calls
## Initial implementation - covers gemini + image gen calls
import json
import time
import uuid
from copy import deepcopy
from functools import partial
@@ -61,6 +62,7 @@ from litellm.types.llms.vertex_ai import (
UsageMetadata,
)
from litellm.types.utils import (
ChatCompletionAudioResponse,
ChatCompletionTokenLogprob,
ChoiceLogprobs,
CompletionTokensDetailsWrapper,
@@ -69,7 +71,7 @@ from litellm.types.utils import (
TopLogprob,
Usage,
)
from litellm.utils import CustomStreamWrapper, ModelResponse, supports_reasoning
from litellm.utils import CustomStreamWrapper, ModelResponse, is_base64_encoded, supports_reasoning
from ....utils import _remove_additional_properties, _remove_strict_from_schema
from ..common_utils import VertexAIError, _build_vertex_schema
@@ -676,14 +678,30 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
) -> Tuple[Optional[str], Optional[str]]:
content_str: Optional[str] = None
reasoning_content_str: Optional[str] = None
for part in parts:
_content_str = ""
if "text" in part:
_content_str += part["text"]
elif "inlineData" in part: # base64 encoded image
_content_str += "data:{};base64,{}".format(
part["inlineData"]["mimeType"], part["inlineData"]["data"]
)
text_content = part["text"]
# Check if text content is audio data URI - if so, exclude from text content
if text_content.startswith("data:audio") and ";base64," in text_content:
try:
if is_base64_encoded(text_content):
media_type, _ = text_content.split("data:")[1].split(";base64,")
if media_type.startswith("audio/"):
continue
except (ValueError, IndexError):
# If parsing fails, treat as regular text
pass
_content_str += text_content
elif "inlineData" in part:
mime_type = part["inlineData"]["mimeType"]
data = part["inlineData"]["data"]
# Check if inline data is audio - if so, exclude from text content
if mime_type.startswith("audio/"):
continue
_content_str += "data:{};base64,{}".format(mime_type, data)
if len(_content_str) > 0:
if part.get("thought") is True:
if reasoning_content_str is None:
@@ -696,6 +714,47 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return content_str, reasoning_content_str
def _extract_audio_response_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[ChatCompletionAudioResponse]:
"""Extract audio response from parts if present"""
for part in parts:
if "text" in part:
text_content = part["text"]
# Check if text content contains audio data URI
if text_content.startswith("data:audio") and ";base64," in text_content:
try:
if is_base64_encoded(text_content):
media_type, audio_data = text_content.split("data:")[1].split(";base64,")
if media_type.startswith("audio/"):
expires_at = int(time.time()) + (24 * 60 * 60)
transcript = "" # Gemini doesn't provide transcript
return ChatCompletionAudioResponse(
data=audio_data,
expires_at=expires_at,
transcript=transcript
)
except (ValueError, IndexError):
pass
elif "inlineData" in part:
mime_type = part["inlineData"]["mimeType"]
data = part["inlineData"]["data"]
if mime_type.startswith("audio/"):
expires_at = int(time.time()) + (24 * 60 * 60)
transcript = "" # Gemini doesn't provide transcript
return ChatCompletionAudioResponse(
data=data,
expires_at=expires_at,
transcript=transcript
)
return None
def _transform_parts(
self,
parts: List[HttpxPartType],
@@ -981,8 +1040,17 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
) = VertexGeminiConfig().get_assistant_content_message(
parts=candidate["content"]["parts"]
)
if content is not None:
audio_response = VertexGeminiConfig()._extract_audio_response_from_parts(
parts=candidate["content"]["parts"]
)
if audio_response is not None:
cast(Dict[str, Any], chat_completion_message)["audio"] = audio_response
chat_completion_message["content"] = None # OpenAI spec
elif content is not None:
chat_completion_message["content"] = content
if reasoning_content is not None:
chat_completion_message["reasoning_content"] = reasoning_content
+13
View File
@@ -176,6 +176,18 @@ class GeminiThinkingConfig(TypedDict, total=False):
GeminiResponseModalities = Literal["TEXT", "IMAGE", "AUDIO", "VIDEO"]
class PrebuiltVoiceConfig(TypedDict):
voiceName: str
class VoiceConfig(TypedDict):
prebuiltVoiceConfig: PrebuiltVoiceConfig
class SpeechConfig(TypedDict, total=False):
voiceConfig: VoiceConfig
class GenerationConfig(TypedDict, total=False):
temperature: float
top_p: float
@@ -252,6 +264,7 @@ class RequestBody(TypedDict, total=False):
safetySettings: List[SafetSettingsConfig]
generationConfig: GenerationConfig
cachedContent: str
speechConfig: SpeechConfig
class CachedContentRequestBody(TypedDict, total=False):
@@ -0,0 +1,221 @@
"""
Test Gemini TTS (Text-to-Speech) functionality
"""
import os
import sys
import pytest
from unittest.mock import patch, MagicMock
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.gemini.chat.transformation import GoogleAIStudioGeminiConfig
from litellm.utils import get_supported_openai_params
class TestGeminiTTSTransformation:
"""Test Gemini TTS transformation functionality"""
def test_gemini_tts_model_detection(self):
"""Test that TTS models are correctly identified"""
config = GoogleAIStudioGeminiConfig()
# Test TTS models
assert config.is_model_gemini_audio_model("gemini-2.5-flash-preview-tts") == True
assert config.is_model_gemini_audio_model("gemini-2.5-pro-preview-tts") == True
# Test non-TTS models
assert config.is_model_gemini_audio_model("gemini-2.5-flash") == False
assert config.is_model_gemini_audio_model("gemini-2.5-pro") == False
assert config.is_model_gemini_audio_model("gpt-4o-audio-preview") == False
def test_gemini_tts_supported_params(self):
"""Test that audio parameter is included for TTS models"""
config = GoogleAIStudioGeminiConfig()
# Test TTS model
params = config.get_supported_openai_params("gemini-2.5-flash-preview-tts")
assert "audio" in params
# Test that other standard params are still included
assert "temperature" in params
assert "max_tokens" in params
assert "modalities" in params
# Test non-TTS model
params_non_tts = config.get_supported_openai_params("gemini-2.5-flash")
assert "audio" not in params_non_tts
def test_gemini_tts_audio_parameter_mapping(self):
"""Test audio parameter mapping for TTS models"""
config = GoogleAIStudioGeminiConfig()
non_default_params = {
"audio": {
"voice": "Kore",
"format": "pcm16"
}
}
optional_params = {}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-2.5-flash-preview-tts",
drop_params=False
)
# Check speech config is created
assert "speechConfig" in result
assert "voiceConfig" in result["speechConfig"]
assert "prebuiltVoiceConfig" in result["speechConfig"]["voiceConfig"]
assert result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
# Check response modalities
assert "responseModalities" in result
assert "AUDIO" in result["responseModalities"]
def test_gemini_tts_audio_parameter_with_existing_modalities(self):
"""Test audio parameter mapping when modalities already exist"""
config = GoogleAIStudioGeminiConfig()
non_default_params = {
"audio": {
"voice": "Puck",
"format": "pcm16"
}
}
optional_params = {
"responseModalities": ["TEXT"]
}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-2.5-flash-preview-tts",
drop_params=False
)
# Check that AUDIO is added to existing modalities
assert "responseModalities" in result
assert "TEXT" in result["responseModalities"]
assert "AUDIO" in result["responseModalities"]
def test_gemini_tts_no_audio_parameter(self):
"""Test that non-audio parameters are handled normally"""
config = GoogleAIStudioGeminiConfig()
non_default_params = {
"temperature": 0.7,
"max_tokens": 100
}
optional_params = {}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-2.5-flash-preview-tts",
drop_params=False
)
# Should not have speech config
assert "speechConfig" not in result
# Should not automatically add audio modalities
assert "responseModalities" not in result
def test_gemini_tts_invalid_audio_parameter(self):
"""Test handling of invalid audio parameter"""
config = GoogleAIStudioGeminiConfig()
non_default_params = {
"audio": "invalid_string" # Should be dict
}
optional_params = {}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-2.5-flash-preview-tts",
drop_params=False
)
# Should not create speech config for invalid audio param
assert "speechConfig" not in result
def test_gemini_tts_empty_audio_parameter(self):
"""Test handling of empty audio parameter"""
config = GoogleAIStudioGeminiConfig()
non_default_params = {
"audio": {}
}
optional_params = {}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-2.5-flash-preview-tts",
drop_params=False
)
# Should still set response modalities even with empty audio config
assert "responseModalities" in result
assert "AUDIO" in result["responseModalities"]
def test_gemini_tts_audio_format_validation(self):
"""Test audio format validation for TTS models"""
config = GoogleAIStudioGeminiConfig()
# Test invalid format
non_default_params = {
"audio": {
"voice": "Kore",
"format": "wav" # Invalid format
}
}
optional_params = {}
with pytest.raises(ValueError, match="Unsupported audio format for Gemini TTS models"):
config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-2.5-flash-preview-tts",
drop_params=False
)
def test_gemini_tts_utils_integration(self):
"""Test integration with LiteLLM utils functions"""
# Test that get_supported_openai_params works with TTS models
params = get_supported_openai_params("gemini-2.5-flash-preview-tts", "gemini")
assert "audio" in params
# Test non-TTS model
params_non_tts = get_supported_openai_params("gemini-2.5-flash", "gemini")
assert "audio" not in params_non_tts
def test_gemini_tts_completion_mock():
"""Test Gemini TTS completion with mocked response"""
with patch('litellm.completion') as mock_completion:
# Mock a successful TTS response
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
mock_response.choices[0].message.content = "Generated audio response"
mock_completion.return_value = mock_response
# Test completion call with audio parameter
response = litellm.completion(
model="gemini-2.5-flash-preview-tts",
messages=[{"role": "user", "content": "Say hello"}],
audio={"voice": "Kore", "format": "pcm16"}
)
assert response is not None
assert response.choices[0].message.content is not None
if __name__ == "__main__":
pytest.main([__file__])