From b4d43c2c9b8414ff94e2ccd26d690d0427c28480 Mon Sep 17 00:00:00 2001 From: javimp2003uma Date: Sun, 16 Aug 2026 11:29:47 +0200 Subject: [PATCH] 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. --- .../elevenlabs/text_to_speech/ws_handler.py | 133 ++++++++ litellm/proxy/proxy_server.py | 140 +++++++++ .../test_elevenlabs_ws_tts_handler.py | 296 ++++++++++++++++++ 3 files changed, 569 insertions(+) create mode 100644 litellm/llms/elevenlabs/text_to_speech/ws_handler.py create mode 100644 tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py diff --git a/litellm/llms/elevenlabs/text_to_speech/ws_handler.py b/litellm/llms/elevenlabs/text_to_speech/ws_handler.py new file mode 100644 index 00000000000..12ad5a0c1e8 --- /dev/null +++ b/litellm/llms/elevenlabs/text_to_speech/ws_handler.py @@ -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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bda6fc25499..3f9f9a572c0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py b/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py new file mode 100644 index 00000000000..6f7621a4c9c --- /dev/null +++ b/tests/test_litellm/llms/elevenlabs/test_elevenlabs_ws_tts_handler.py @@ -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