mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat: Add ModelsLab text-to-speech provider
- Implements ModelsLabTextToSpeechConfig following BaseTextToSpeechConfig pattern
- Key-in-body authentication (MODELSLAB_API_KEY env var)
- Async polling: processing status → poll /voice/fetch/{id} until success
- Downloads audio from output URL and returns HttpxBinaryResponseContent
- Voice mappings: alloy/echo/fable/onyx/nova/shimmer → ModelsLab voice IDs
- Language support: english, spanish, french, german, italian, and more
- Adds MODELSLAB to LlmProviders enum and get_provider_text_to_speech_config()
- 9 unit tests (all mocked, no real network calls)
This commit is contained in:
parent
98974771fd
commit
1de97cf2e1
8 changed files with 403 additions and 0 deletions
0
litellm/llms/modelslab/__init__.py
Normal file
0
litellm/llms/modelslab/__init__.py
Normal file
0
litellm/llms/modelslab/text_to_speech/__init__.py
Normal file
0
litellm/llms/modelslab/text_to_speech/__init__.py
Normal file
232
litellm/llms/modelslab/text_to_speech/transformation.py
Normal file
232
litellm/llms/modelslab/text_to_speech/transformation.py
Normal file
|
|
@ -0,0 +1,232 @@
|
|||
"""
|
||||
ModelsLab Text-to-Speech transformation for LiteLLM.
|
||||
|
||||
NOTE: ModelsLab uses key-in-body authentication. The MODELSLAB_API_KEY
|
||||
will appear in the request body. Handle accordingly.
|
||||
|
||||
ModelsLab TTS API: https://docs.modelslab.com
|
||||
"""
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.text_to_speech.transformation import (
|
||||
BaseTextToSpeechConfig,
|
||||
TextToSpeechRequestData,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
HttpxBinaryResponseContent = Any
|
||||
|
||||
MODELSLAB_TTS_URL = "https://modelslab.com/api/v6/voice/text_to_speech"
|
||||
MODELSLAB_TTS_FETCH_URL = "https://modelslab.com/api/v6/voice/fetch/{request_id}"
|
||||
MODELSLAB_POLL_INTERVAL = 5
|
||||
MODELSLAB_POLL_TIMEOUT = 300
|
||||
|
||||
# OpenAI voice → ModelsLab voice_id mappings
|
||||
VOICE_MAPPINGS: Dict[str, int] = {
|
||||
"alloy": 1, # neutral
|
||||
"echo": 2, # male
|
||||
"fable": 3, # warm
|
||||
"onyx": 4, # deep male
|
||||
"nova": 5, # female
|
||||
"shimmer": 6, # clear female
|
||||
}
|
||||
|
||||
LANGUAGE_MAPPINGS: Dict[str, str] = {
|
||||
"en": "english",
|
||||
"es": "spanish",
|
||||
"fr": "french",
|
||||
"de": "german",
|
||||
"it": "italian",
|
||||
"pt": "portuguese",
|
||||
"zh": "chinese",
|
||||
"ja": "japanese",
|
||||
"ko": "korean",
|
||||
"hi": "hindi",
|
||||
"ar": "arabic",
|
||||
}
|
||||
|
||||
|
||||
class ModelsLabTextToSpeechConfig(BaseTextToSpeechConfig):
|
||||
"""
|
||||
Configuration for ModelsLab Text-to-Speech.
|
||||
|
||||
ModelsLab TTS uses an async pattern:
|
||||
1. POST /api/v6/voice/text_to_speech → {status: success, output: "url"} or {status: processing, request_id: ...}
|
||||
2. If processing, poll POST /api/v6/voice/fetch/{request_id} with {key} body until done
|
||||
3. Download audio from output URL
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._api_key: Optional[str] = None
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
return ["voice", "response_format", "speed"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
model: str,
|
||||
optional_params: Dict,
|
||||
voice: Optional[Union[str, Dict]] = None,
|
||||
drop_params: bool = False,
|
||||
kwargs: Dict = {},
|
||||
) -> Tuple[Optional[str], Dict]:
|
||||
mapped: Dict[str, Any] = {}
|
||||
|
||||
# Resolve voice_id
|
||||
voice_id: Optional[int] = None
|
||||
if isinstance(voice, str) and voice.strip():
|
||||
voice_id = VOICE_MAPPINGS.get(voice.lower(), 1)
|
||||
elif isinstance(voice, dict):
|
||||
voice_id = voice.get("voice_id", 1)
|
||||
if voice_id:
|
||||
mapped["voice_id"] = voice_id
|
||||
|
||||
# Map speed (ModelsLab accepts 0.5-2.0)
|
||||
if "speed" in optional_params:
|
||||
try:
|
||||
mapped["speed"] = float(optional_params["speed"])
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# Language from extra params
|
||||
if "language" in kwargs:
|
||||
lang = kwargs["language"]
|
||||
mapped["language"] = LANGUAGE_MAPPINGS.get(lang, lang)
|
||||
|
||||
return voice, mapped
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""Key-in-body auth — only Content-Type goes in headers."""
|
||||
api_key = api_key or litellm.api_key or get_secret_str("MODELSLAB_API_KEY")
|
||||
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"ModelsLab API key is required. Set MODELSLAB_API_KEY or pass api_key."
|
||||
)
|
||||
|
||||
self._api_key = api_key
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
model: str,
|
||||
api_base: Optional[str],
|
||||
litellm_params: dict,
|
||||
) -> str:
|
||||
if api_base:
|
||||
return api_base.rstrip("/")
|
||||
return MODELSLAB_TTS_URL
|
||||
|
||||
def transform_text_to_speech_request(
|
||||
self,
|
||||
model: str,
|
||||
input: str,
|
||||
voice: Optional[str],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
headers: dict,
|
||||
) -> TextToSpeechRequestData:
|
||||
body: Dict[str, Any] = {
|
||||
"key": self._api_key,
|
||||
"prompt": input,
|
||||
"language": optional_params.pop("language", "english"),
|
||||
"voice_id": optional_params.pop("voice_id", VOICE_MAPPINGS.get(voice or "", 1)),
|
||||
"speed": optional_params.pop("speed", 1.0),
|
||||
}
|
||||
# Pass through any remaining provider-specific params
|
||||
body.update(optional_params)
|
||||
return TextToSpeechRequestData(dict_body=body)
|
||||
|
||||
def transform_text_to_speech_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> "HttpxBinaryResponseContent":
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
response_data = raw_response.json()
|
||||
status = response_data.get("status", "")
|
||||
request_id = str(response_data.get("request_id", ""))
|
||||
|
||||
if status == "error":
|
||||
raise BaseLLMException(
|
||||
status_code=raw_response.status_code,
|
||||
message=response_data.get("message", "ModelsLab TTS failed"),
|
||||
headers=dict(raw_response.headers),
|
||||
)
|
||||
|
||||
if status == "processing":
|
||||
response_data = self._poll_tts_sync(request_id)
|
||||
|
||||
audio_url = response_data.get("output", "")
|
||||
if not audio_url:
|
||||
raise BaseLLMException(
|
||||
status_code=500,
|
||||
message="ModelsLab TTS returned no audio URL",
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Download the audio file
|
||||
client: HTTPHandler = _get_httpx_client()
|
||||
audio_response = client.get(audio_url)
|
||||
audio_response.raise_for_status()
|
||||
|
||||
return HttpxBinaryResponseContent(audio_response)
|
||||
|
||||
def _poll_tts_sync(
|
||||
self,
|
||||
request_id: str,
|
||||
timeout: int = MODELSLAB_POLL_TIMEOUT,
|
||||
interval: int = MODELSLAB_POLL_INTERVAL,
|
||||
) -> Dict:
|
||||
"""Poll the ModelsLab TTS fetch endpoint until done."""
|
||||
fetch_url = MODELSLAB_TTS_FETCH_URL.format(request_id=request_id)
|
||||
body = {"key": self._api_key}
|
||||
client: HTTPHandler = _get_httpx_client()
|
||||
deadline = time.time() + timeout
|
||||
|
||||
while time.time() < deadline:
|
||||
time.sleep(interval)
|
||||
resp = client.post(fetch_url, json=body)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
if data.get("status") in ("success", "error"):
|
||||
return data
|
||||
|
||||
raise BaseLLMException(
|
||||
status_code=408,
|
||||
message=f"ModelsLab TTS timed out after {timeout}s (request_id={request_id})",
|
||||
headers={},
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Dict,
|
||||
) -> BaseLLMException:
|
||||
raise BaseLLMException(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -3099,6 +3099,7 @@ class LlmProviders(str, Enum):
|
|||
DEEPINFRA = "deepinfra"
|
||||
PERPLEXITY = "perplexity"
|
||||
MISTRAL = "mistral"
|
||||
MODELSLAB = "modelslab"
|
||||
MILVUS = "milvus"
|
||||
GROQ = "groq"
|
||||
A2A = "a2a"
|
||||
|
|
|
|||
|
|
@ -8897,6 +8897,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return AWSPollyTextToSpeechConfig()
|
||||
elif litellm.LlmProviders.MODELSLAB == provider:
|
||||
from litellm.llms.modelslab.text_to_speech.transformation import (
|
||||
ModelsLabTextToSpeechConfig,
|
||||
)
|
||||
|
||||
return ModelsLabTextToSpeechConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
0
tests/test_litellm/llms/modelslab/__init__.py
Normal file
0
tests/test_litellm/llms/modelslab/__init__.py
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
"""
|
||||
Tests for ModelsLab TTS transformation. All mocked — no real network calls.
|
||||
"""
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.modelslab.text_to_speech.transformation import (
|
||||
ModelsLabTextToSpeechConfig,
|
||||
VOICE_MAPPINGS,
|
||||
)
|
||||
|
||||
|
||||
class TestModelsLabTTSTransformation:
|
||||
|
||||
def setup_method(self):
|
||||
self.config = ModelsLabTextToSpeechConfig()
|
||||
self.config._api_key = "test-api-key"
|
||||
self.mock_logging = Mock()
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# validate_environment
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
def test_validate_environment_no_auth_header(self):
|
||||
"""Key-in-body auth: only Content-Type, no Authorization header."""
|
||||
with patch(
|
||||
"litellm.llms.modelslab.text_to_speech.transformation.get_secret_str",
|
||||
return_value="test-key",
|
||||
):
|
||||
headers = self.config.validate_environment(headers={}, model="default")
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
assert "Authorization" not in headers
|
||||
|
||||
def test_validate_environment_raises_without_key(self):
|
||||
with patch(
|
||||
"litellm.llms.modelslab.text_to_speech.transformation.get_secret_str",
|
||||
return_value=None,
|
||||
):
|
||||
with pytest.raises(ValueError, match="MODELSLAB_API_KEY"):
|
||||
self.config.validate_environment(headers={}, model="default")
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# map_openai_params / voice mapping
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
def test_map_openai_params_voice_mapping(self):
|
||||
"""OpenAI voice names map to ModelsLab voice_id integers."""
|
||||
_, params = self.config.map_openai_params(
|
||||
model="default", optional_params={}, voice="nova"
|
||||
)
|
||||
assert params["voice_id"] == VOICE_MAPPINGS["nova"]
|
||||
|
||||
def test_map_openai_params_speed(self):
|
||||
_, params = self.config.map_openai_params(
|
||||
model="default", optional_params={"speed": "1.5"}, voice="alloy"
|
||||
)
|
||||
assert params["speed"] == 1.5
|
||||
|
||||
def test_get_supported_openai_params(self):
|
||||
params = self.config.get_supported_openai_params("default")
|
||||
assert "voice" in params
|
||||
assert "response_format" in params
|
||||
assert "speed" in params
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# transform_text_to_speech_request
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
def test_transform_request_key_in_body(self):
|
||||
"""API key must be in body, not headers."""
|
||||
result = self.config.transform_text_to_speech_request(
|
||||
model="default",
|
||||
input="Hello world",
|
||||
voice="alloy",
|
||||
optional_params={"language": "english", "voice_id": 1, "speed": 1.0},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
body = result["dict_body"]
|
||||
assert body["key"] == "test-api-key"
|
||||
assert body["prompt"] == "Hello world"
|
||||
assert "language" in body
|
||||
|
||||
def test_transform_request_voice_default(self):
|
||||
"""Unknown voice defaults to voice_id 1."""
|
||||
result = self.config.transform_text_to_speech_request(
|
||||
model="default",
|
||||
input="Test",
|
||||
voice="unknown_voice",
|
||||
optional_params={"language": "english", "speed": 1.0},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
body = result["dict_body"]
|
||||
assert body["voice_id"] == 1 # default
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# transform_text_to_speech_response
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
def test_transform_response_success_downloads_audio(self):
|
||||
"""Success response fetches audio URL and returns HttpxBinaryResponseContent."""
|
||||
mock_resp = Mock(spec=httpx.Response)
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = {}
|
||||
mock_resp.json.return_value = {
|
||||
"status": "success",
|
||||
"output": "https://cdn.modelslab.com/output/audio.mp3",
|
||||
}
|
||||
|
||||
mock_audio_resp = Mock(spec=httpx.Response)
|
||||
mock_audio_resp.status_code = 200
|
||||
mock_audio_resp.content = b"fake-audio-bytes"
|
||||
|
||||
with patch.object(self.config, "_poll_tts_sync") as mock_poll, \
|
||||
patch(
|
||||
"litellm.llms.modelslab.text_to_speech.transformation._get_httpx_client"
|
||||
) as mock_client:
|
||||
mock_client.return_value.get.return_value = mock_audio_resp
|
||||
result = self.config.transform_text_to_speech_response(
|
||||
model="default",
|
||||
raw_response=mock_resp,
|
||||
logging_obj=self.mock_logging,
|
||||
)
|
||||
mock_poll.assert_not_called() # no polling needed for success
|
||||
|
||||
assert result is not None
|
||||
|
||||
def test_transform_response_processing_polls(self):
|
||||
"""Processing response triggers polling then downloads audio."""
|
||||
mock_resp = Mock(spec=httpx.Response)
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = {}
|
||||
mock_resp.json.return_value = {
|
||||
"status": "processing",
|
||||
"request_id": "req_abc",
|
||||
"eta": 5,
|
||||
}
|
||||
|
||||
poll_result = {
|
||||
"status": "success",
|
||||
"output": "https://cdn.modelslab.com/output/audio2.mp3",
|
||||
}
|
||||
|
||||
mock_audio_resp = Mock(spec=httpx.Response)
|
||||
mock_audio_resp.status_code = 200
|
||||
mock_audio_resp.content = b"fake-audio-bytes-2"
|
||||
|
||||
with patch.object(
|
||||
self.config, "_poll_tts_sync", return_value=poll_result
|
||||
) as mock_poll, patch(
|
||||
"litellm.llms.modelslab.text_to_speech.transformation._get_httpx_client"
|
||||
) as mock_client:
|
||||
mock_client.return_value.get.return_value = mock_audio_resp
|
||||
result = self.config.transform_text_to_speech_response(
|
||||
model="default",
|
||||
raw_response=mock_resp,
|
||||
logging_obj=self.mock_logging,
|
||||
)
|
||||
mock_poll.assert_called_once_with("req_abc")
|
||||
|
||||
assert result is not None
|
||||
Loading…
Add table
Reference in a new issue