mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(elevenlabs): add WebSocket streaming-input TTS endpoint to proxy
Adds support for ElevenLabs' /v1/text-to-speech/{voice_id}/stream-input
WebSocket API, which enables low-latency TTS when text arrives token by
token (e.g. from an LLM output stream).
New proxy WebSocket route: /v1/audio/speech/stream-input
- Authenticates via LiteLLM key, connects upstream with xi-api-key
- Sends BOS automatically; client streams {"text":"..."} chunks and
closes with {"text":""}, matching the ElevenLabs wire protocol
- Forwards audio JSON (base64 + alignment fields) back to client as-is
- Computes per-character cost from model_prices_and_context_window.json
and fires standard aspeech logging/callback chain on session close
New module: litellm/llms/elevenlabs/text_to_speech/ws_handler.py
- build_elevenlabs_ws_url: constructs wss:// URL with path-encoded
voice_id and model/format query params
- stream_input_tts: sends BOS, runs bidirectional relay, returns total
character count for cost tracking
- Pydantic TypeAdapters for both client and server message shapes;
no Any in public signatures
15 unit tests covering URL construction, relay character counting,
EOS detection, flush/try_trigger_generation passthrough, binary frame
handling, and BOS voice/generation-config injection.
This commit is contained in:
parent
13d94ec546
commit
b4d43c2c9b
3 changed files with 569 additions and 0 deletions
133
litellm/llms/elevenlabs/text_to_speech/ws_handler.py
Normal file
133
litellm/llms/elevenlabs/text_to_speech/ws_handler.py
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Final, Required
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.websockets import WebSocket
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
_WS_BASE_URL: Final = "wss://api.elevenlabs.io"
|
||||
_WS_PATH: Final = "/v1/text-to-speech/{voice_id}/stream-input"
|
||||
|
||||
|
||||
class VoiceSettings(TypedDict, total=False):
|
||||
stability: ReadOnly[float]
|
||||
similarity_boost: ReadOnly[float]
|
||||
style: ReadOnly[float]
|
||||
use_speaker_boost: ReadOnly[bool]
|
||||
speed: ReadOnly[float]
|
||||
|
||||
|
||||
class GenerationConfig(TypedDict, total=False):
|
||||
chunk_length_schedule: ReadOnly[list[float]]
|
||||
|
||||
|
||||
class _TtsClientMessage(TypedDict, total=False):
|
||||
text: Required[ReadOnly[str]]
|
||||
flush: ReadOnly[bool]
|
||||
try_trigger_generation: ReadOnly[bool]
|
||||
voice_settings: ReadOnly[VoiceSettings]
|
||||
generator_config: ReadOnly[GenerationConfig]
|
||||
|
||||
|
||||
class _TtsServerMessage(TypedDict, total=False):
|
||||
audio: ReadOnly[str]
|
||||
isFinal: ReadOnly[bool] # camelCase matches ElevenLabs API field name
|
||||
|
||||
|
||||
_CLIENT_MSG_ADAPTER: Final = TypeAdapter(_TtsClientMessage)
|
||||
_SERVER_MSG_ADAPTER: Final = TypeAdapter(_TtsServerMessage)
|
||||
|
||||
|
||||
def build_elevenlabs_ws_url(
|
||||
model: str,
|
||||
voice_id: str,
|
||||
output_format: str,
|
||||
api_base: str | None = None,
|
||||
) -> str:
|
||||
raw_base: Final = (api_base or get_secret_str("ELEVENLABS_API_BASE") or _WS_BASE_URL).rstrip("/")
|
||||
ws_base: Final = raw_base.replace("https://", "wss://").replace("http://", "ws://")
|
||||
encoded_voice: Final = encode_url_path_segment(voice_id, field_name="voice_id")
|
||||
path: Final = _WS_PATH.format(voice_id=encoded_voice)
|
||||
return f"{ws_base}{path}?model_id={model}&output_format={output_format}"
|
||||
|
||||
|
||||
async def _relay_client_to_upstream(
|
||||
client_ws: WebSocket,
|
||||
upstream: ClientConnection,
|
||||
) -> tuple[int, ...]:
|
||||
chunk_lengths: list[int] = [] # mutable-ok: local accumulator, converted to immutable tuple on return
|
||||
async for raw in client_ws.iter_text():
|
||||
msg = _CLIENT_MSG_ADAPTER.validate_json(raw)
|
||||
await upstream.send(json.dumps(dict(msg)))
|
||||
text = msg.get("text", "")
|
||||
chunk_lengths.append(len(text))
|
||||
if not text:
|
||||
break
|
||||
return tuple(chunk_lengths)
|
||||
|
||||
|
||||
async def _relay_upstream_to_client(
|
||||
upstream: ClientConnection,
|
||||
client_ws: WebSocket,
|
||||
) -> None:
|
||||
async for raw in upstream:
|
||||
payload = raw if isinstance(raw, str) else raw.decode()
|
||||
await client_ws.send_text(payload)
|
||||
msg = _SERVER_MSG_ADAPTER.validate_json(payload)
|
||||
if msg.get("isFinal"):
|
||||
break
|
||||
|
||||
|
||||
async def stream_input_tts(
|
||||
*,
|
||||
client_ws: WebSocket,
|
||||
model: str,
|
||||
voice_id: str,
|
||||
output_format: str = "mp3_44100_128",
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
voice_settings: VoiceSettings | None = None,
|
||||
generation_config: GenerationConfig | None = None,
|
||||
) -> int:
|
||||
"""
|
||||
Relay a streaming-input TTS session between a client WebSocket and ElevenLabs.
|
||||
|
||||
The proxy sends BOS automatically on connect using the provided voice/generation
|
||||
settings. The client then sends text chunks as {"text": "..."} and signals end of
|
||||
stream with {"text": ""}. Audio responses (JSON with base64 audio) are forwarded
|
||||
back to the client as-is.
|
||||
|
||||
Returns the total number of text characters sent (used for per-character cost tracking).
|
||||
"""
|
||||
import websockets
|
||||
|
||||
key: Final = api_key or get_secret_str("ELEVENLABS_API_KEY")
|
||||
if key is None:
|
||||
raise ValueError("ElevenLabs API key is required. Set ELEVENLABS_API_KEY.")
|
||||
|
||||
url: Final = build_elevenlabs_ws_url(model, voice_id, output_format, api_base)
|
||||
|
||||
bos: dict[str, object] = {"text": " "} # mutable-ok: fields added conditionally before first send
|
||||
if voice_settings is not None:
|
||||
bos["voice_settings"] = voice_settings
|
||||
if generation_config is not None:
|
||||
bos["generation_config"] = generation_config
|
||||
|
||||
async with websockets.connect(url, additional_headers={"xi-api-key": key}) as upstream:
|
||||
await upstream.send(json.dumps(bos))
|
||||
results: Final = await asyncio.gather(
|
||||
_relay_client_to_upstream(client_ws, upstream),
|
||||
_relay_upstream_to_client(upstream, client_ws),
|
||||
)
|
||||
|
||||
chunk_lengths: Final = results[0]
|
||||
return sum(chunk_lengths)
|
||||
|
|
@ -10866,6 +10866,146 @@ async def realtime_websocket_endpoint(
|
|||
await websocket.close(code=1011, reason="Internal server error")
|
||||
|
||||
|
||||
######################################################################
|
||||
|
||||
# /v1/audio/speech/stream-input Endpoint
|
||||
|
||||
######################################################################
|
||||
|
||||
|
||||
@app.websocket("/v1/audio/speech/stream-input")
|
||||
@app.websocket("/audio/speech/stream-input")
|
||||
async def elevenlabs_tts_stream_input_endpoint(
|
||||
websocket: WebSocket,
|
||||
model: str = fastapi.Query(
|
||||
...,
|
||||
description="Model ID, e.g. 'elevenlabs/eleven_multilingual_v2' or 'eleven_multilingual_v2'.",
|
||||
),
|
||||
voice: str = fastapi.Query(
|
||||
...,
|
||||
description="ElevenLabs voice ID or an OpenAI voice-name alias (alloy, coral, …).",
|
||||
),
|
||||
output_format: str = fastapi.Query(
|
||||
"mp3_44100_128",
|
||||
description="Audio output format accepted by ElevenLabs (e.g. 'pcm_44100', 'mp3_44100_128').",
|
||||
),
|
||||
stability: float | None = fastapi.Query(None, description="Voice stability (0-1)."),
|
||||
similarity_boost: float | None = fastapi.Query(None, description="Voice similarity boost (0-1)."),
|
||||
style: float | None = fastapi.Query(None, description="Voice style (0-1, v2+ models only)."),
|
||||
speed: float | None = fastapi.Query(None, description="Speaking speed (0.7-1.2)."),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket),
|
||||
):
|
||||
"""
|
||||
ElevenLabs WebSocket streaming-input TTS endpoint.
|
||||
|
||||
The proxy sends BOS automatically. The client streams text chunks as
|
||||
{"text": "..."} messages and signals end-of-stream with {"text": ""}.
|
||||
Audio responses are forwarded back as JSON (same format as ElevenLabs).
|
||||
|
||||
Query parameters map directly to ElevenLabs voice/model options; no
|
||||
per-message auth is needed because the proxy injects xi-api-key upstream.
|
||||
"""
|
||||
from datetime import datetime
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.elevenlabs.text_to_speech.transformation import (
|
||||
ElevenLabsTextToSpeechConfig,
|
||||
)
|
||||
from litellm.llms.elevenlabs.text_to_speech.ws_handler import (
|
||||
VoiceSettings,
|
||||
stream_input_tts,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
elevenlabs_model: Final = model.removeprefix("elevenlabs/")
|
||||
litellm_model: Final = f"elevenlabs/{elevenlabs_model}"
|
||||
|
||||
try:
|
||||
await can_key_call_resolved_model(
|
||||
model=litellm_model,
|
||||
llm_model_list=llm_model_list,
|
||||
valid_token=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyException as e:
|
||||
await websocket.close(code=1008, reason=e.message[:120])
|
||||
return
|
||||
|
||||
await websocket.accept()
|
||||
|
||||
elevenlabs_config: Final = ElevenLabsTextToSpeechConfig()
|
||||
voice_id: Final = elevenlabs_config._extract_voice_id(voice)
|
||||
|
||||
raw_voice_settings: Final = {
|
||||
k: v
|
||||
for k, v in {
|
||||
"stability": stability,
|
||||
"similarity_boost": similarity_boost,
|
||||
"style": style,
|
||||
"speed": speed,
|
||||
}.items()
|
||||
if v is not None
|
||||
}
|
||||
voice_settings: Final[VoiceSettings | None] = cast(VoiceSettings, raw_voice_settings) if raw_voice_settings else None
|
||||
|
||||
start_time: Final = datetime.now()
|
||||
litellm_call_id: Final = str(uuid4())
|
||||
|
||||
logging_obj: Final = Logging(
|
||||
model=litellm_model,
|
||||
messages=[],
|
||||
stream=True,
|
||||
call_type="aspeech",
|
||||
start_time=start_time,
|
||||
litellm_call_id=litellm_call_id,
|
||||
function_id="elevenlabs_tts_stream_input",
|
||||
)
|
||||
|
||||
try:
|
||||
total_chars: Final = await stream_input_tts(
|
||||
client_ws=websocket,
|
||||
model=elevenlabs_model,
|
||||
voice_id=voice_id,
|
||||
output_format=output_format,
|
||||
voice_settings=voice_settings,
|
||||
)
|
||||
|
||||
end_time: Final = datetime.now()
|
||||
|
||||
try:
|
||||
model_info: Final = litellm.get_model_info(
|
||||
model=litellm_model, custom_llm_provider="elevenlabs"
|
||||
)
|
||||
cost_per_char: Final = model_info.get("input_cost_per_character") or 0.0
|
||||
except Exception:
|
||||
cost_per_char = 0.0
|
||||
|
||||
response_cost: Final = total_chars * cost_per_char
|
||||
|
||||
mock_http_response: Final = httpx.Response(200, content=b"")
|
||||
result: Final = HttpxBinaryResponseContent(mock_http_response)
|
||||
result._hidden_params = {
|
||||
"response_cost": response_cost,
|
||||
"model_id": litellm_model,
|
||||
"litellm_call_id": litellm_call_id,
|
||||
}
|
||||
await logging_obj.async_success_handler(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("ElevenLabs TTS stream-input error")
|
||||
try:
|
||||
await websocket.close(code=1011, reason="Internal server error")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
######################################################################
|
||||
|
||||
# /v1/assistant Endpoints
|
||||
|
|
|
|||
|
|
@ -0,0 +1,296 @@
|
|||
"""
|
||||
Unit tests for litellm/llms/elevenlabs/text_to_speech/ws_handler.py.
|
||||
|
||||
These tests verify the WebSocket URL builder and the relay helpers without
|
||||
establishing real network connections.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.elevenlabs.text_to_speech.ws_handler import (
|
||||
_relay_client_to_upstream,
|
||||
_relay_upstream_to_client,
|
||||
build_elevenlabs_ws_url,
|
||||
stream_input_tts,
|
||||
)
|
||||
|
||||
|
||||
class TestBuildElevenLabsWsUrl:
|
||||
def test_default_base_url(self) -> None:
|
||||
url = build_elevenlabs_ws_url(
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="21m00Tcm4TlvDq8ikWAM",
|
||||
output_format="mp3_44100_128",
|
||||
)
|
||||
assert url.startswith("wss://api.elevenlabs.io/v1/text-to-speech/")
|
||||
assert "?model_id=eleven_multilingual_v2&output_format=mp3_44100_128" in url
|
||||
assert "21m00Tcm4TlvDq8ikWAM" in url
|
||||
|
||||
def test_custom_api_base_https_converted_to_wss(self) -> None:
|
||||
url = build_elevenlabs_ws_url(
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="voice123",
|
||||
output_format="pcm_44100",
|
||||
api_base="https://custom.elevenlabs.io",
|
||||
)
|
||||
assert url.startswith("wss://custom.elevenlabs.io/")
|
||||
assert "http" not in url.split("?")[0]
|
||||
|
||||
def test_custom_api_base_http_converted_to_ws(self) -> None:
|
||||
url = build_elevenlabs_ws_url(
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="voice123",
|
||||
output_format="pcm_44100",
|
||||
api_base="http://localhost:8080",
|
||||
)
|
||||
assert url.startswith("ws://localhost:8080/")
|
||||
|
||||
def test_voice_id_with_special_chars_is_encoded(self) -> None:
|
||||
url = build_elevenlabs_ws_url(
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="voice with spaces",
|
||||
output_format="mp3_44100_128",
|
||||
)
|
||||
assert " " not in url
|
||||
assert "voice%20with%20spaces" in url
|
||||
|
||||
def test_query_params_include_model_and_format(self) -> None:
|
||||
url = build_elevenlabs_ws_url(
|
||||
model="eleven_turbo_v2",
|
||||
voice_id="abc",
|
||||
output_format="ulaw_8000",
|
||||
)
|
||||
assert "model_id=eleven_turbo_v2" in url
|
||||
assert "output_format=ulaw_8000" in url
|
||||
|
||||
def test_trailing_slash_stripped_from_base(self) -> None:
|
||||
url = build_elevenlabs_ws_url(
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="abc",
|
||||
output_format="mp3_44100_128",
|
||||
api_base="wss://api.elevenlabs.io/",
|
||||
)
|
||||
assert "//" not in url.replace("wss://", "")
|
||||
|
||||
|
||||
class TestRelayClientToUpstream:
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_text_chunks_and_counts_chars(self) -> None:
|
||||
chunks = [
|
||||
json.dumps({"text": "Hello, "}),
|
||||
json.dumps({"text": "world! "}),
|
||||
json.dumps({"text": ""}),
|
||||
]
|
||||
client_ws = MagicMock()
|
||||
client_ws.iter_text = self._make_iter(chunks)
|
||||
|
||||
upstream = AsyncMock()
|
||||
upstream.send = AsyncMock()
|
||||
|
||||
result = await _relay_client_to_upstream(client_ws, upstream)
|
||||
|
||||
assert result == (7, 7, 0)
|
||||
assert upstream.send.call_count == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stops_on_empty_text_eos(self) -> None:
|
||||
chunks = [
|
||||
json.dumps({"text": "First "}),
|
||||
json.dumps({"text": ""}),
|
||||
json.dumps({"text": "Should not be sent"}),
|
||||
]
|
||||
client_ws = MagicMock()
|
||||
client_ws.iter_text = self._make_iter(chunks)
|
||||
|
||||
upstream = AsyncMock()
|
||||
upstream.send = AsyncMock()
|
||||
|
||||
result = await _relay_client_to_upstream(client_ws, upstream)
|
||||
|
||||
assert sum(result) == 6
|
||||
assert upstream.send.call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passes_flush_and_try_trigger_fields(self) -> None:
|
||||
msg = {"text": "Go! ", "flush": True, "try_trigger_generation": True}
|
||||
chunks = [json.dumps(msg), json.dumps({"text": ""})]
|
||||
client_ws = MagicMock()
|
||||
client_ws.iter_text = self._make_iter(chunks)
|
||||
|
||||
upstream = AsyncMock()
|
||||
sent_payloads: list[dict[str, Any]] = []
|
||||
|
||||
async def capture_send(payload: str) -> None:
|
||||
sent_payloads.append(json.loads(payload))
|
||||
|
||||
upstream.send = capture_send
|
||||
|
||||
await _relay_client_to_upstream(client_ws, upstream)
|
||||
|
||||
assert sent_payloads[0]["flush"] is True
|
||||
assert sent_payloads[0]["try_trigger_generation"] is True
|
||||
|
||||
@staticmethod
|
||||
def _make_iter(items: list[str]):
|
||||
async def _gen():
|
||||
for item in items:
|
||||
yield item
|
||||
|
||||
return lambda: _gen()
|
||||
|
||||
|
||||
class TestRelayUpstreamToClient:
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_audio_messages_to_client(self) -> None:
|
||||
messages = [
|
||||
json.dumps({"audio": "base64audio1"}),
|
||||
json.dumps({"audio": "base64audio2"}),
|
||||
json.dumps({"isFinal": True}),
|
||||
]
|
||||
upstream = MagicMock()
|
||||
upstream.__aiter__ = self._make_iter(messages)
|
||||
|
||||
client_ws = AsyncMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
|
||||
await _relay_upstream_to_client(upstream, client_ws)
|
||||
|
||||
assert client_ws.send_text.call_count == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stops_after_is_final(self) -> None:
|
||||
messages = [
|
||||
json.dumps({"isFinal": True}),
|
||||
json.dumps({"audio": "should_not_be_forwarded"}),
|
||||
]
|
||||
upstream = MagicMock()
|
||||
upstream.__aiter__ = self._make_iter(messages)
|
||||
|
||||
client_ws = AsyncMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
|
||||
await _relay_upstream_to_client(upstream, client_ws)
|
||||
|
||||
assert client_ws.send_text.call_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handles_binary_upstream_messages(self) -> None:
|
||||
messages = [b'{"isFinal": true}']
|
||||
upstream = MagicMock()
|
||||
upstream.__aiter__ = self._make_iter(messages)
|
||||
|
||||
client_ws = AsyncMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
|
||||
await _relay_upstream_to_client(upstream, client_ws)
|
||||
|
||||
client_ws.send_text.assert_called_once_with('{"isFinal": true}')
|
||||
|
||||
@staticmethod
|
||||
def _make_iter(items: list[str | bytes]):
|
||||
async def _gen(self):
|
||||
for item in items:
|
||||
yield item
|
||||
|
||||
return _gen
|
||||
|
||||
|
||||
def _make_upstream_mock(messages: list[str], send_capture: list[dict[str, Any]] | None = None):
|
||||
"""Build a minimal async context manager that acts as a websockets connection."""
|
||||
|
||||
class _FakeUpstream:
|
||||
def __aiter__(self):
|
||||
return self._gen()
|
||||
|
||||
async def _gen(self):
|
||||
for msg in messages:
|
||||
yield msg
|
||||
|
||||
async def send(self, payload: str) -> None:
|
||||
if send_capture is not None:
|
||||
send_capture.append(json.loads(payload))
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
pass
|
||||
|
||||
return _FakeUpstream()
|
||||
|
||||
|
||||
class TestStreamInputTts:
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_total_char_count(self) -> None:
|
||||
client_ws = AsyncMock()
|
||||
client_ws.iter_text = lambda: self._aiter([
|
||||
json.dumps({"text": "Hello, "}),
|
||||
json.dumps({"text": "world! "}),
|
||||
json.dumps({"text": ""}),
|
||||
])
|
||||
|
||||
upstream_messages = [
|
||||
json.dumps({"audio": "dGVzdA=="}),
|
||||
json.dumps({"isFinal": True}),
|
||||
]
|
||||
|
||||
with (
|
||||
patch("litellm.llms.elevenlabs.text_to_speech.ws_handler.get_secret_str", return_value="test-key"),
|
||||
patch("websockets.connect", return_value=_make_upstream_mock(upstream_messages)),
|
||||
):
|
||||
total = await stream_input_tts(
|
||||
client_ws=client_ws,
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="21m00Tcm4TlvDq8ikWAM",
|
||||
output_format="mp3_44100_128",
|
||||
)
|
||||
|
||||
assert total == 14
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sends_bos_with_voice_settings(self) -> None:
|
||||
client_ws = AsyncMock()
|
||||
client_ws.iter_text = lambda: self._aiter([json.dumps({"text": ""})])
|
||||
|
||||
sent_messages: list[dict[str, Any]] = []
|
||||
upstream = _make_upstream_mock([json.dumps({"isFinal": True})], send_capture=sent_messages)
|
||||
|
||||
with (
|
||||
patch("litellm.llms.elevenlabs.text_to_speech.ws_handler.get_secret_str", return_value="test-key"),
|
||||
patch("websockets.connect", return_value=upstream),
|
||||
):
|
||||
await stream_input_tts(
|
||||
client_ws=client_ws,
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="abc",
|
||||
output_format="mp3_44100_128",
|
||||
voice_settings={"stability": 0.5, "similarity_boost": 0.75},
|
||||
generation_config={"chunk_length_schedule": [120]},
|
||||
)
|
||||
|
||||
bos = sent_messages[0]
|
||||
assert bos["text"] == " "
|
||||
assert bos["voice_settings"] == {"stability": 0.5, "similarity_boost": 0.75}
|
||||
assert bos["generation_config"] == {"chunk_length_schedule": [120]}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_when_api_key_missing(self) -> None:
|
||||
client_ws = MagicMock()
|
||||
|
||||
with patch("litellm.llms.elevenlabs.text_to_speech.ws_handler.get_secret_str", return_value=None):
|
||||
with pytest.raises(ValueError, match="ELEVENLABS_API_KEY"):
|
||||
await stream_input_tts(
|
||||
client_ws=client_ws,
|
||||
model="eleven_multilingual_v2",
|
||||
voice_id="abc",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _aiter(items: list[str]):
|
||||
for item in items:
|
||||
yield item
|
||||
Loading…
Add table
Reference in a new issue